7std::optional<u64> dispatch_float_width(
nat_t width,
auto f) {
10 case i: return f.template operator()<i>();
17std::optional<u64> dispatch_int_width(
nat_t width,
auto f) {
20 case i: return f.template operator()<i>();
28std::optional<u64> fold_float_unary_bits(
u64 a, std::invocable<
w2f<w>>
auto f) {
30 auto x = fe::bitcast_resize<T>(
a);
31 return fe::bitcast_resize<u64>(
static_cast<T
>(
f(x)));
37 auto x = fe::bitcast_resize<T>(
a);
38 auto y = fe::bitcast_resize<T>(b);
39 return fe::bitcast_resize<u64>(
static_cast<T
>(
f(x, y)));
43constexpr long double signed_min() {
47 return static_cast<long double>(std::numeric_limits<w2s<w>>::min());
51constexpr long double signed_max() {
55 return static_cast<long double>(std::numeric_limits<w2s<w>>::max());
59constexpr long double unsigned_max() {
60 return static_cast<long double>(std::numeric_limits<w2u<w>>::max());
64std::optional<u64> encode_signed(
long double x) {
65 if constexpr (
w == 1) {
66 if (x == -1.0L)
return 1_u64;
67 if (x == 0.0L)
return 0_u64;
70 return fe::bitcast_resize<u64>(
static_cast<w2s<w>>(x));
75std::optional<u64> encode_unsigned(
long double x) {
76 return fe::bitcast_resize<u64>(
static_cast<w2u<w>>(x));
80long double decode_signed(
u64 a) {
82 return fe::bitcast_resize<bool>(
a) ? -1.0L : 0.0L;
84 return static_cast<long double>(fe::bitcast_resize<w2s<w>>(
a));
88long double decode_unsigned(
u64 a) {
89 return static_cast<long double>(fe::bitcast_resize<w2u<w>>(
a));
93std::optional<u64> fold_float_to_signed_bits(std::floating_point
auto x) {
94 if (!std::isfinite(x))
return {};
96 auto truncated = std::trunc(
static_cast<long double>(x));
97 if (truncated < signed_min<w>() || truncated > signed_max<w>())
return {};
98 return encode_signed<w>(truncated);
102std::optional<u64> fold_float_to_unsigned_bits(std::floating_point
auto x) {
103 if (!std::isfinite(x))
return {};
105 auto truncated = std::trunc(
static_cast<long double>(x));
106 if (truncated < 0.0L || truncated > unsigned_max<w>())
return {};
107 return encode_unsigned<w>(truncated);
111template<
class Id, Id
id, nat_t w>
112std::optional<u64> fold_unary_lit(
u64 a) {
113 if constexpr (std::is_same_v<Id, tri>) {
114 if constexpr (
false) {}
115 else if constexpr (
id ==
tri:: sin )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std:: sin (x); });
116 else if constexpr (
id ==
tri:: cos )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std:: cos (x); });
117 else if constexpr (
id ==
tri:: tan )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std:: tan (x); });
118 else if constexpr (
id ==
tri:: sinh)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std:: sinh(x); });
119 else if constexpr (
id ==
tri:: cosh)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std:: cosh(x); });
120 else if constexpr (
id ==
tri:: tanh)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std:: tanh(x); });
121 else if constexpr (
id ==
tri::asin )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::asin (x); });
122 else if constexpr (
id ==
tri::acos )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::acos (x); });
123 else if constexpr (
id ==
tri::atan )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::atan (x); });
124 else if constexpr (
id ==
tri::asinh)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::asinh(x); });
125 else if constexpr (
id ==
tri::acosh)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::acosh(x); });
126 else if constexpr (
id ==
tri::atanh)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::atanh(x); });
127 else fe::unreachable();
128 }
else if constexpr (std::is_same_v<Id, rt>) {
129 if constexpr (
false) {}
130 else if constexpr (
id ==
rt::sq)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::sqrt(x); });
131 else if constexpr (
id ==
rt::cb)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::cbrt(x); });
132 else static_assert(
false,
"missing sub tag");
133 }
else if constexpr (std::is_same_v<Id, exp>) {
134 if constexpr (
false) {}
135 else if constexpr (
id ==
exp::exp)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::exp(x); });
136 else if constexpr (
id ==
exp::exp2)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::exp2(x); });
137 else if constexpr (
id ==
exp::exp10)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::pow(
decltype(x)(10), x); });
138 else if constexpr (
id ==
exp::log)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::log(x); });
139 else if constexpr (
id ==
exp::log2)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::log2(x); });
140 else if constexpr (
id ==
exp::log10)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::log10(x); });
141 else fe::unreachable();
142 }
else if constexpr (std::is_same_v<Id, er>) {
143 if constexpr (
false) {}
144 else if constexpr (
id ==
er::f )
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::erf (x); });
145 else if constexpr (
id ==
er::fc)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::erfc(x); });
146 else static_assert(
false,
"missing sub tag");
147 }
else if constexpr (std::is_same_v<Id, gamma>) {
148 if constexpr (
false) {}
149 else if constexpr (
id ==
gamma::t)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::tgamma(x); });
150 else if constexpr (
id ==
gamma::l)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::lgamma(x); });
151 else static_assert(
false,
"missing sub tag");
152 }
else if constexpr (std::is_same_v<Id, round>) {
153 if constexpr (
false) {}
154 else if constexpr (
id ==
round::f)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::floor(x); });
155 else if constexpr (
id ==
round::c)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::ceil(x); });
156 else if constexpr (
id ==
round::r)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::round(x); });
157 else if constexpr (
id ==
round::t)
return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::trunc(x); });
158 else static_assert(
false,
"missing sub tag");
160 static_assert(
false,
"missing tag");
165template<
class Id, nat_t w>
166std::optional<u64> fold_unary_lit(
u64 a) {
167 if constexpr (std::is_same_v<Id, abs>)
168 return fold_float_unary_bits<w>(
a, [](
auto x) {
return std::abs(x); });
170 static_assert(
false,
"missing tag");
174const Def*
fold(World& world,
const Def* type,
const Def*
a) {
175 if (
a->isa<
Bot>())
return world.bot(type);
178 if (
auto width =
isa_f(
a->type()))
179 if (
auto res = dispatch_float_width(*width, [&]<
nat_t w>() {
return fold_unary_lit<Id, w>(*la); }))
180 return world.lit(type, *res);
185template<
class Id, Id
id, nat_t w>
186std::optional<u64> fold_binary_lit(
u64 a,
u64 b) {
188 auto x = fe::bitcast_resize<T>(
a);
189 auto y = fe::bitcast_resize<T>(b);
191 if constexpr (std::is_same_v<Id, arith>) {
193 if constexpr (
false) {}
194 else if constexpr (
id ==
arith::add)
return fe::bitcast_resize<u64>(
static_cast<T
>(x + y));
195 else if constexpr (
id ==
arith::sub)
return fe::bitcast_resize<u64>(
static_cast<T
>(x - y));
196 else if constexpr (
id ==
arith::mul)
return fe::bitcast_resize<u64>(
static_cast<T
>(x * y));
197 else if constexpr (
id ==
arith::div)
return fe::bitcast_resize<u64>(
static_cast<T
>(x / y));
198 else if constexpr (
id ==
arith::rem)
return fe::bitcast_resize<u64>(
static_cast<T
>(
rem(x, y)));
199 else static_assert(
false,
"missing sub tag");
201 }
else if constexpr (std::is_same_v<Id, math::extrema>) {
204 if (x == T(-0.0) && y == T(+0.0)) {
206 }
else if (x == T(+0.0) && y == T(-0.0)) {
213 else if (std::isnan(y))
218 static_assert(
false,
"missing sub tag");
221 return fe::bitcast_resize<u64>(result);
222 }
else if constexpr (std::is_same_v<Id, pow>) {
223 return fe::bitcast_resize<u64>(
static_cast<T
>(std::pow(x, y)));
224 }
else if constexpr (std::is_same_v<Id, cmp>) {
225 using std::isunordered;
227 result |= ((
id &
cmp::u) !=
cmp::f) && isunordered(x, y);
233 static_assert(
false,
"missing tag");
237template<
class Id, Id
id>
238const Def*
fold(World& world,
const Def* type,
const Def*
a) {
239 if (
a->isa<
Bot>())
return world.bot(type);
242 if (
auto width =
isa_f(
a->type()))
243 if (
auto res = dispatch_float_width(*width, [&]<
nat_t w>() {
return fold_unary_lit<Id, id, w>(*la); }))
244 return world.lit(type, *res);
250template<
class Id, Id
id>
251const Def*
fold(World& world,
const Def* type,
const Def*&
a,
const Def*& b) {
252 if (
a->isa<
Bot>() || b->isa<
Bot>())
return world.bot(type);
256 if (
auto width =
isa_f(
a->type()))
258 = dispatch_float_width(*width, [&]<
nat_t w>() {
return fold_binary_lit<Id, id, w>(*la, *lb); }))
259 return world.lit(type, *res);
277const Def* reassociate(Id
id, World& world, [[maybe_unused]]
const App* ab,
const Def*
a,
const Def* b) {
282 auto la =
a->isa<
Lit>();
283 auto [x, y] = xy ? xy->template args<2>() : std::array<const Def*, 2>{
nullptr,
nullptr};
284 auto [z,
w] = zw ? zw->template args<2>() : std::array<const Def*, 2>{
nullptr,
nullptr};
290 auto check_mode = [&](
const App* app) {
291 auto app_m =
Lit::isa(app->decurry()->arg());
292 if (!app_m || !fe::has_flag(
static_cast<Mode>(*app_m),
Mode::reassoc))
return false;
297 if (!check_mode(ab))
return nullptr;
298 if (lx && !check_mode(xy->decurry()))
return nullptr;
299 if (lz && !check_mode(zw->decurry()))
return nullptr;
301 auto make_op = [&](
const Def*
a,
const Def* b) {
return world.call(
id,
mode,
Defs{
a, b}); };
303 if (la && lz)
return make_op(make_op(
a, z), w);
304 if (lx && lz)
return make_op(make_op(x, z), make_op(y, w));
305 if (lz)
return make_op(z, make_op(
a, w));
306 if (lx)
return make_op(x, make_op(y, b));
310template<conv
id, nat_t sw, nat_t dw>
311std::optional<u64> fold_conv_lit(
u64 a) {
315 if constexpr (
false) {}
316 else if constexpr (
id ==
conv::s2f)
return fe::bitcast_resize<u64>(
static_cast<D>(decode_signed<sw>(
a)));
317 else if constexpr (
id ==
conv::u2f)
return fe::bitcast_resize<u64>(
static_cast<D>(decode_unsigned<sw>(
a)));
318 else if constexpr (
id ==
conv::f2s)
return fold_float_to_signed_bits<dw>(fe::bitcast_resize<S>(
a));
319 else if constexpr (
id ==
conv::f2u)
return fold_float_to_unsigned_bits<dw>(fe::bitcast_resize<S>(
a));
320 else if constexpr (
id ==
conv::f2f)
return fe::bitcast_resize<u64>(
static_cast<D>(fe::bitcast_resize<S>(
a)));
321 else static_assert(
false,
"missing sub tag");
325template<conv
id, nat_t sw>
326std::optional<u64> fold_conv_dst(
nat_t dw,
u64 a) {
328 return dispatch_float_width(dw, [&]<
nat_t d>() {
return fold_conv_lit<id, sw, d>(
a); });
330 return dispatch_int_width(dw, [&]<
nat_t d>() {
return fold_conv_lit<id, sw, d>(
a); });
336 return dispatch_int_width(sw, [&]<
nat_t s>() {
return fold_conv_dst<id, s>(dw,
a); });
338 return dispatch_float_width(sw, [&]<
nat_t s>() {
return fold_conv_dst<id, s>(dw,
a); });
345 auto& world = type->world();
346 auto callee =
c->as<
App>();
347 auto [a, b] = arg->
projs<2>();
348 auto mode = callee->decurry()->arg();
350 auto w =
isa_f(a->type());
352 if (
auto result = fold<arith, id>(world, type, a, b))
return result;
357 auto zero =
lit_f(world, *w, 0.0);
358 auto one =
lit_f(world, *w, 1.0);
359 auto two =
lit_f(world, *w, 2.0);
361 if (
auto la = a->isa<
Lit>()) {
362 if (zero && la == zero) {
372 if (one && la == one) {
383 if (
auto lb = b->isa<
Lit>()) {
384 if (zero && lb == zero) {
389 default: fe::unreachable();
398 case arith::sub:
if (zero)
return zero;
break;
400 case arith::div:
if (one )
return one ;
break;
407 if (
auto res = reassociate<arith>(
id, world, callee, a, b))
return res;
409 return world.raw_app(type, callee, {a, b});
414 auto& world = type->world();
415 auto callee =
c->as<
App>();
416 auto [a, b] = arg->
projs<2>();
417 auto m = callee->decurry()->arg();
421 if (
auto lit = fold<extrema, id>(world, type, a, b))
return lit;
431 return world.raw_app(type,
c, {a, b});
436 auto& world = type->world();
437 if (
auto lit = fold<tri, id>(world, type, arg))
return lit;
442 auto& world = type->world();
443 auto [a, b] = arg->
projs<2>();
444 if (
auto lit = fold<
pow,
pow(0)>(world, type, a, b))
return lit;
450 auto& world = type->world();
451 if (
auto lit = fold<rt, id>(world, type, arg))
return lit;
457 auto& world = type->world();
458 if (
auto lit = fold<exp, id>(world, type, arg))
return lit;
464 auto& world = type->world();
465 if (
auto lit = fold<er, id>(world, type, arg))
return lit;
471 auto& world = type->world();
472 if (
auto lit = fold<gamma, id>(world, type, arg))
return lit;
478 auto& world = type->world();
479 auto callee =
c->as<
App>();
480 auto [a, b] = arg->
projs<2>();
482 if (
auto result = fold<cmp, id>(world, type, a, b))
return result;
483 if (
id ==
cmp::f)
return world.lit_ff();
484 if (
id ==
cmp::t)
return world.lit_tt();
486 return world.raw_app(type, callee, {a, b});
491 auto& world = dst_t->
world();
492 auto s_t = x->
type()->as<
App>();
493 auto d_t = dst_t->as<
App>();
499 if (s_t == d_t)
return x;
500 if (x->isa<
Bot>())
return world.bot(d_t);
508 if (
auto l =
Lit::isa(x); l && sw && dw)
509 if (
auto res = fold_conv<id>(*sw, *dw, *l))
return world.lit(d_t, *res);
515 auto& world = type->world();
516 if (
auto lit = fold<abs>(world, type, arg))
return lit;
522 auto& world = type->world();
523 if (
auto lit = fold<round, id>(world, type, arg))
return lit;
static auto isa(const Def *def)
World & world() const noexcept
auto projs(Projector auto f) const
Splits this Def via Def::projections into an Array (if A == std::dynamic_extent) or std::array (other...
const Def * type() const noexcept
Yields the "raw" type of this Def (maybe nullptr).
static bool greater(const Def *a, const Def *b)
static constexpr nat_t size2bitwidth(nat_t n)
static std::optional< T > isa(const Def *def)
#define MIM_math_NORMALIZER_IMPL
const Def * normalize_extrema(const Def *type, const Def *c, const Def *arg)
const Def * normalize_er(const Def *type, const Def *, const Def *arg)
const Def * normalize_cmp(const Def *type, const Def *c, const Def *arg)
const Def * normalize_abs(const Def *type, const Def *, const Def *arg)
const Def * normalize_gamma(const Def *type, const Def *, const Def *arg)
Mode
Allowed optimizations for a specific operation.
@ reassoc
Allow reassociation transformations for floating-point operations.
@ bot
Alias for Mode::fast.
const Lit * lit_f(World &w, std::floating_point auto val)
const Def * normalize_arith(const Def *type, const Def *c, const Def *arg)
const Def * normalize_round(const Def *type, const Def *, const Def *arg)
std::optional< nat_t > isa_f(const Def *def)
const Def * normalize_tri(const Def *type, const Def *, const Def *arg)
const Def * normalize_exp(const Def *type, const Def *, const Def *arg)
const Def * normalize_rt(const Def *type, const Def *, const Def *arg)
const Def * normalize_pow(const Def *type, const Def *, const Def *arg)
const Def * normalize_conv(const Def *dst_t, const Def *, const Def *x)
typename detail::w2f_< w >::type w2f
fe::View< const Def * > Defs
constexpr bool is_commutative(Id)
typename detail::w2s_< w >::type w2s
constexpr bool is_associative(Id id)
typename detail::w2u_< w >::type w2u
#define MIM_1_8_16_32_64(X)