native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
constexpr_float.h
1// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
2#pragma once
3#include <array>
4#include <bit>
5#include <cstddef>
6#include <cstdint>
7#include <type_traits>
8
9namespace native::detail::constexpr_float {
10 // Integer encodings keep constant evaluation independent of the host FP
11 // environment. This implementation does not model exception status or traps.
12 template<unsigned ExponentBits,unsigned FractionBits>
13 struct binary_format {
14 static_assert(ExponentBits>=2 && ExponentBits<=11 && FractionBits>=1 &&
15 FractionBits<=52 && 1+ExponentBits+FractionBits<=64);
16 static constexpr unsigned exponent_bits=ExponentBits;
17 static constexpr unsigned fraction_bits=FractionBits;
18 static constexpr unsigned width=1+ExponentBits+FractionBits;
19 using bits_type=std::conditional_t<(width<=16),std::uint16_t,
20 std::conditional_t<(width<=32),std::uint32_t,std::uint64_t>>;
21 static constexpr std::uint64_t fraction_mask=(std::uint64_t{1}<<FractionBits)-1;
22 static constexpr std::uint64_t exponent_max=(std::uint64_t{1}<<ExponentBits)-1;
23 static constexpr std::uint64_t exponent_mask=exponent_max<<FractionBits;
24 static constexpr std::uint64_t sign_mask=std::uint64_t{1}<<(width-1);
25 static constexpr std::uint64_t quiet_mask=std::uint64_t{1}<<(FractionBits-1);
26 static constexpr int bias=(1<<(ExponentBits-1))-1;
27 static constexpr int min_normal_exponent=1-bias;
28 static constexpr int min_subnormal_exponent=min_normal_exponent-int(FractionBits);
29 static constexpr int max_normal_exponent=bias;
30 };
31 using binary16=binary_format<5,10>;
32 using bfloat16=binary_format<8,7>;
33 using binary32=binary_format<8,23>;
34 using binary64=binary_format<11,52>;
35
36 enum class rounding { nearest_even, downward, upward, toward_zero, to_odd };
37 enum class nan_propagation { signaling_first, first, default_nan };
38 enum class fma_nan_order { addend_first, multiplicands_first };
39
40 struct policy {
41 nan_propagation nan=nan_propagation::signaling_first;
42 fma_nan_order fma_order=fma_nan_order::addend_first;
43 bool flush_inputs=false;
44 // Flush before rounding, including tiny values that would round to normal.
45 bool flush_outputs=false;
46 // Legacy BFloat16 dot arithmetic overflows to infinity even with round-odd.
47 bool odd_overflow_infinity=false;
48 bool default_nan_negative=false;
49 // Arm AH=0: a quiet addend does not mask an invalid 0 * infinity product.
50 bool invalid_product_overrides_quiet_addend=true;
51 };
52
53 template<class F> constexpr bool is_nan(typename F::bits_type x) noexcept {
54 return (x&F::exponent_mask)==F::exponent_mask && (x&F::fraction_mask)!=0;
55 }
56 template<class F> constexpr bool is_signaling_nan(typename F::bits_type x) noexcept {
57 return is_nan<F>(x) && !(x&F::quiet_mask);
58 }
59 template<class F> constexpr bool is_infinite(typename F::bits_type x) noexcept {
60 return (x&~F::sign_mask)==F::exponent_mask;
61 }
62 template<class F> constexpr bool is_zero(typename F::bits_type x) noexcept {
63 return (x&~F::sign_mask)==0;
64 }
65 template<class F> constexpr typename F::bits_type quiet_nan(typename F::bits_type x) noexcept {
66 return typename F::bits_type(x|F::quiet_mask);
67 }
68 template<class F> constexpr typename F::bits_type default_nan(policy p={}) noexcept {
69 return typename F::bits_type(F::exponent_mask|F::quiet_mask|
70 (p.default_nan_negative?F::sign_mask:0));
71 }
72 template<class F> constexpr typename F::bits_type flush_input(
73 typename F::bits_type x,policy p) noexcept {
74 return p.flush_inputs && !(x&F::exponent_mask)
75 ?typename F::bits_type(x&F::sign_mask):x;
76 }
77
78 template<class F,std::size_t N>
79 constexpr typename F::bits_type select_nan(
80 std::array<typename F::bits_type,N> const & operands,policy p) noexcept {
81 if(p.nan==nan_propagation::default_nan) return default_nan<F>(p);
82 if(p.nan==nan_propagation::signaling_first)
83 for(auto x:operands) if(is_signaling_nan<F>(x)) return quiet_nan<F>(x);
84 for(auto x:operands) if(is_nan<F>(x)) return quiet_nan<F>(x);
85 return default_nan<F>(p);
86 }
87
88 struct finite_parts { bool sign; std::uint64_t significand; int exponent; };
89 template<class F> constexpr finite_parts unpack(typename F::bits_type x) noexcept {
90 auto e=unsigned((x&F::exponent_mask)>>F::fraction_bits);
91 return {(x&F::sign_mask)!=0,
92 (x&F::fraction_mask)|(e?std::uint64_t{1}<<F::fraction_bits:0),
93 e?int(e)-F::bias-int(F::fraction_bits):F::min_subnormal_exponent};
94 }
95
96 // Fixed-capacity unsigned magnitude. Arithmetic uses only uint64_t, including
97 // on Windows targets without a language-level 128-bit integer type.
98 template<std::size_t N> struct magnitude {
99 std::array<std::uint64_t,N> words{};
100 constexpr unsigned bit_width() const noexcept {
101 for(std::size_t i=N;i>0;--i)
102 if(words[i-1]) return unsigned((i-1)*64+std::bit_width(words[i-1]));
103 return 0;
104 }
105 constexpr void insert(std::uint64_t x,unsigned shift) noexcept {
106 auto i=shift/64;
107 auto n=shift%64;
108 if(i<N) words[i]|=x<<n;
109 if(n && i+1<N) words[i+1]|=x>>(64-n);
110 }
111 constexpr std::uint64_t extract(unsigned shift) const noexcept {
112 auto i=shift/64;
113 auto n=shift%64;
114 auto x=i<N?words[i]>>n:0;
115 if(n && i+1<N) x|=words[i+1]<<(64-n);
116 return x;
117 }
118 constexpr bool below(unsigned bit) const noexcept {
119 auto full=bit/64;
120 for(std::size_t i=0;i<N && i<full;++i) if(words[i]) return true;
121 return full<N && bit%64 && (words[full]&((std::uint64_t{1}<<(bit%64))-1));
122 }
123 constexpr int compare(magnitude const & other) const noexcept {
124 for(std::size_t i=N;i>0;--i)
125 if(words[i-1]!=other.words[i-1]) return words[i-1]>other.words[i-1]?1:-1;
126 return 0;
127 }
128 constexpr void add(magnitude const & other) noexcept {
129 bool carry=false;
130 for(std::size_t i=0;i<N;++i) {
131 auto x=words[i],sum=x+other.words[i];
132 auto next=sum+std::uint64_t(carry);
133 carry=sum<x || next<sum;
134 words[i]=next;
135 }
136 }
137 // Precondition: *this >= other.
138 constexpr void subtract(magnitude const & other) noexcept {
139 bool borrow=false;
140 for(std::size_t i=0;i<N;++i) {
141 auto x=words[i],difference=x-other.words[i];
142 auto next=difference-std::uint64_t(borrow);
143 borrow=x<other.words[i] || difference<std::uint64_t(borrow);
144 words[i]=next;
145 }
146 }
147 };
148
149 struct double_word { std::uint64_t low,high; };
150 constexpr double_word multiply(std::uint64_t a,std::uint64_t b) noexcept {
151 constexpr auto mask=std::uint64_t{0xffffffff};
152 auto a0=a&mask,a1=a>>32,b0=b&mask,b1=b>>32;
153 auto low=a0*b0;
154 auto middle=a1*b0+(low>>32);
155 auto high=middle>>32;
156 middle=(middle&mask)+a0*b1;
157 return {(middle<<32)|(low&mask),a1*b1+high+(middle>>32)};
158 }
159
160 template<class F> constexpr typename F::bits_type overflow(
161 bool sign,rounding mode,policy p) noexcept {
162 bool infinity=mode==rounding::nearest_even ||
163 (mode==rounding::upward && !sign) || (mode==rounding::downward && sign) ||
164 (mode==rounding::to_odd && p.odd_overflow_infinity);
165 return typename F::bits_type((sign?F::sign_mask:0)|
166 (infinity?F::exponent_mask:F::exponent_mask-1));
167 }
168
169 // Round sign * value * 2^exponent once to F. All discarded bits remain
170 // available, even after cancellation of a product and a distant addend.
171 template<class F,std::size_t N>
172 constexpr typename F::bits_type round_pack(bool sign,magnitude<N> const & value,
173 int exponent,rounding mode=rounding::nearest_even,policy p={}) noexcept {
174 auto sign_bits=sign?F::sign_mask:0;
175 auto width=value.bit_width();
176 if(!width) return typename F::bits_type(sign_bits);
177 auto highest=exponent+int(width)-1;
178 if(p.flush_outputs && highest<F::min_normal_exponent)
179 return typename F::bits_type(sign_bits);
180 auto quantum=highest-int(F::fraction_bits);
181 if(quantum<F::min_subnormal_exponent) quantum=F::min_subnormal_exponent;
182 auto shift=quantum-exponent;
183 auto mantissa=shift>=0?value.extract(unsigned(shift)):value.words[0]<<unsigned(-shift);
184 bool guard=shift>0 && (value.extract(unsigned(shift-1))&1);
185 bool sticky=shift>1 && value.below(unsigned(shift-1));
186 bool inexact=guard || sticky;
187 bool increment=(mode==rounding::nearest_even && guard && (sticky || (mantissa&1))) ||
188 (mode==rounding::downward && sign && inexact) ||
189 (mode==rounding::upward && !sign && inexact);
190 if(mode==rounding::to_odd && inexact) mantissa|=1;
191 else if(increment) ++mantissa;
192 if(mantissa>=(std::uint64_t{1}<<(F::fraction_bits+1))) {mantissa>>=1;++quantum;}
193 auto encoded_exponent=mantissa>=(std::uint64_t{1}<<F::fraction_bits)
194 ?quantum+int(F::fraction_bits)+F::bias:0;
195 if(encoded_exponent>=int(F::exponent_max)) return overflow<F>(sign,mode,p);
196 return typename F::bits_type(sign_bits|(std::uint64_t(encoded_exponent)<<F::fraction_bits)|
197 (mantissa&F::fraction_mask));
198 }
199
200 // Resize a NaN payload. Mixed-precision arithmetic can retain a signaling
201 // operand until the operation selects among all its NaNs; FP conversions
202 // themselves always request quieting. Precondition: x is a NaN encoding.
203 template<class To,class From>
204 constexpr typename To::bits_type resize_nan(typename From::bits_type x,
205 bool quiet=true) noexcept {
206 auto payload=std::uint64_t(x&From::fraction_mask);
207 if constexpr(To::fraction_bits>=From::fraction_bits)
208 payload<<=To::fraction_bits-From::fraction_bits;
209 else payload>>=From::fraction_bits-To::fraction_bits;
210 if(quiet) payload|=To::quiet_mask;
211 else if(!payload) payload=1;
212 return typename To::bits_type(((x&From::sign_mask)?To::sign_mask:0)|To::exponent_mask|payload);
213 }
214
215 template<class To,class From>
216 constexpr typename To::bits_type convert_bits(typename From::bits_type x,
217 rounding mode=rounding::nearest_even,policy p={}) noexcept {
218 auto sign=(x&From::sign_mask)!=0;
219 auto sign_bits=sign?To::sign_mask:0;
220 if(is_nan<From>(x)) {
221 if(p.nan==nan_propagation::default_nan) return default_nan<To>(p);
222 return resize_nan<To,From>(x);
223 }
224 if(is_infinite<From>(x)) return typename To::bits_type(sign_bits|To::exponent_mask);
225 auto parts=unpack<From>(flush_input<From>(x,p));
226 magnitude<1> value{{parts.significand}};
227 return round_pack<To>(sign,value,parts.exponent,mode,p);
228 }
229
230 // Capacity includes the full exponent span of two products plus a carry.
231 template<class F> inline constexpr std::size_t arithmetic_words=
232 (2*(F::max_normal_exponent-F::min_subnormal_exponent)+3+63)/64;
233
234 template<class F,std::size_t N>
235 constexpr typename F::bits_type sum_magnitudes(magnitude<N> a,bool sign_a,
236 magnitude<N> b,bool sign_b,int exponent,rounding mode,policy p) noexcept {
237 if(sign_a==sign_b) {
238 a.add(b);
239 return round_pack<F>(sign_a,a,exponent,mode,p);
240 }
241 auto order=a.compare(b);
242 if(!order) return typename F::bits_type(mode==rounding::downward?F::sign_mask:0);
243 if(order>0) {a.subtract(b);return round_pack<F>(sign_a,a,exponent,mode,p);}
244 b.subtract(a);
245 return round_pack<F>(sign_b,b,exponent,mode,p);
246 }
247
248 template<class F>
249 constexpr typename F::bits_type add_bits(typename F::bits_type x,typename F::bits_type y,
250 rounding mode=rounding::nearest_even,policy p={}) noexcept {
251 x=flush_input<F>(x,p);y=flush_input<F>(y,p);
252 if(is_nan<F>(x) || is_nan<F>(y)) return select_nan<F>(std::array{x,y},p);
253 bool ix=is_infinite<F>(x),iy=is_infinite<F>(y);
254 if(ix && iy && ((x^y)&F::sign_mask)) return default_nan<F>(p);
255 if(ix || iy) return ix?x:y;
256 auto a=unpack<F>(x),b=unpack<F>(y);
257 auto exponent=a.exponent<b.exponent?a.exponent:b.exponent;
258 magnitude<arithmetic_words<F>> av{},bv{};
259 av.insert(a.significand,unsigned(a.exponent-exponent));
260 bv.insert(b.significand,unsigned(b.exponent-exponent));
261 return sum_magnitudes<F>(av,a.sign,bv,b.sign,exponent,mode,p);
262 }
263
264 template<class F>
265 constexpr typename F::bits_type mul_bits(typename F::bits_type x,typename F::bits_type y,
266 rounding mode=rounding::nearest_even,policy p={}) noexcept {
267 x=flush_input<F>(x,p);y=flush_input<F>(y,p);
268 if(is_nan<F>(x) || is_nan<F>(y)) return select_nan<F>(std::array{x,y},p);
269 bool ix=is_infinite<F>(x),iy=is_infinite<F>(y);
270 if((ix && is_zero<F>(y)) || (iy && is_zero<F>(x))) return default_nan<F>(p);
271 auto a=unpack<F>(x),b=unpack<F>(y);
272 bool sign=a.sign!=b.sign;
273 if(ix || iy) return typename F::bits_type((sign?F::sign_mask:0)|F::exponent_mask);
274 auto product=multiply(a.significand,b.significand);
275 magnitude<2> value{{product.low,product.high}};
276 return round_pack<F>(sign,value,a.exponent+b.exponent,mode,p);
277 }
278
279 // Computes x*y+addend with one rounding. NaN order is explicit: the default
280 // examines addend,x,y for signaling NaNs, then the same order for quiet NaNs.
281 template<class F>
282 constexpr typename F::bits_type fma_bits(typename F::bits_type x,typename F::bits_type y,
283 typename F::bits_type addend,rounding mode=rounding::nearest_even,policy p={}) noexcept {
284 x=flush_input<F>(x,p);y=flush_input<F>(y,p);addend=flush_input<F>(addend,p);
285 bool ix=is_infinite<F>(x),iy=is_infinite<F>(y),iz=is_infinite<F>(addend);
286 bool invalid=(ix && is_zero<F>(y)) || (iy && is_zero<F>(x));
287 if(is_nan<F>(x) || is_nan<F>(y) || is_nan<F>(addend)) {
288 if(invalid && is_nan<F>(addend) && !is_signaling_nan<F>(addend) &&
289 p.invalid_product_overrides_quiet_addend) return default_nan<F>(p);
290 return select_nan<F>(p.fma_order==fma_nan_order::addend_first
291 ?std::array{addend,x,y}:std::array{x,y,addend},p);
292 }
293 auto a=unpack<F>(x),b=unpack<F>(y),c=unpack<F>(addend);
294 bool sign=a.sign!=b.sign;
295 if(invalid || ((ix || iy) && iz && sign!=c.sign)) return default_nan<F>(p);
296 if(ix || iy) return typename F::bits_type((sign?F::sign_mask:0)|F::exponent_mask);
297 if(iz) return addend;
298 auto exponent_product=a.exponent+b.exponent;
299 auto exponent=exponent_product<c.exponent?exponent_product:c.exponent;
300 auto product=multiply(a.significand,b.significand);
301 magnitude<arithmetic_words<F>> pv{},cv{};
302 pv.insert(product.low,unsigned(exponent_product-exponent));
303 pv.insert(product.high,unsigned(exponent_product-exponent)+64);
304 cv.insert(c.significand,unsigned(c.exponent-exponent));
305 return sum_magnitudes<F>(pv,sign,cv,c.sign,exponent,mode,p);
306 }
307
308 template<class F>
309 constexpr typename F::bits_type sub_bits(typename F::bits_type x,typename F::bits_type y,
310 rounding mode=rounding::nearest_even,policy p={}) noexcept {
311 // Subtraction does not negate a propagated NaN's sign or payload.
312 if(is_nan<F>(x) || is_nan<F>(y)) return select_nan<F>(std::array{x,y},p);
313 return add_bits<F>(x,typename F::bits_type(y^F::sign_mask),mode,p);
314 }
315
316 template<class F>
317 constexpr typename F::bits_type div_bits(typename F::bits_type x,typename F::bits_type y,
318 rounding mode=rounding::nearest_even,policy p={}) noexcept {
319 x=flush_input<F>(x,p);y=flush_input<F>(y,p);
320 if(is_nan<F>(x) || is_nan<F>(y)) return select_nan<F>(std::array{x,y},p);
321 bool sign=((x^y)&F::sign_mask)!=0;
322 auto sign_bits=sign?F::sign_mask:0;
323 bool ix=is_infinite<F>(x),iy=is_infinite<F>(y),zx=is_zero<F>(x),zy=is_zero<F>(y);
324 if((ix && iy) || (zx && zy)) return default_nan<F>(p);
325 if(ix || zy) return typename F::bits_type(sign_bits|F::exponent_mask);
326 if(iy || zx) return typename F::bits_type(sign_bits);
327 auto a=unpack<F>(x),b=unpack<F>(y);
328 auto sa=F::fraction_bits+1-unsigned(std::bit_width(a.significand));
329 auto sb=F::fraction_bits+1-unsigned(std::bit_width(b.significand));
330 auto remainder=a.significand<<sa,divisor=b.significand<<sb;
331 std::uint64_t quotient=0;
332 // Normalized operands give a quotient in [1/2,2). Retain at least two
333 // discarded bits, then jam the exact remainder into the low sticky bit.
334 constexpr unsigned digits=F::fraction_bits+4;
335 for(unsigned bit=0;bit<digits;++bit) {
336 quotient<<=1;
337 if(remainder>=divisor) {remainder-=divisor;quotient|=1;}
338 remainder<<=1;
339 }
340 if(remainder) quotient|=1;
341 magnitude<1> value{{quotient}};
342 return round_pack<F>(sign,value,a.exponent-int(sa)-b.exponent+int(sb)-int(digits-1),mode,p);
343 }
344
345 template<class F>
346 constexpr typename F::bits_type sqrt_bits(typename F::bits_type x,
347 rounding mode=rounding::nearest_even,policy p={}) noexcept {
348 x=flush_input<F>(x,p);
349 if(is_nan<F>(x)) return select_nan<F>(std::array{x},p);
350 if(is_zero<F>(x)) return x;
351 if(x&F::sign_mask) return default_nan<F>(p);
352 if(is_infinite<F>(x)) return x;
353 auto a=unpack<F>(x);
354 auto highest=a.exponent+int(std::bit_width(a.significand))-1;
355 // Floor division also handles negative odd exponents.
356 auto root_exponent=(highest-(highest<0 && highest%2!=0))/2;
357 auto quantum=root_exponent-int(F::fraction_bits)-2;
358 magnitude<2> radicand{};
359 radicand.insert(a.significand,unsigned(a.exponent-2*quantum));
360 std::uint64_t root=0,remainder=0;
361 // Restoring square root, consuming two radicand bits per step. With at
362 // most 55 root bits even binary64's remainder fits comfortably in uint64_t.
363 for(unsigned pair=(radicand.bit_width()+1)/2;pair>0;--pair) {
364 remainder=(remainder<<2)|(radicand.extract(2*(pair-1))&3);
365 root<<=1;
366 auto trial=(root<<1)|1;
367 if(remainder>=trial) {remainder-=trial;root|=1;}
368 }
369 if(remainder) root|=1;
370 return round_pack<F>(false,magnitude<1>{{root}},quantum,mode,p);
371 }
372
373 template<class F>
374 constexpr typename F::bits_type round_integral_bits(typename F::bits_type x,
375 rounding mode=rounding::nearest_even,policy p={}) noexcept {
376 x=flush_input<F>(x,p);
377 if(is_nan<F>(x)) return select_nan<F>(std::array{x},p);
378 if(is_infinite<F>(x)) return x;
379 auto a=unpack<F>(x);
380 if(a.exponent>=0 || !a.significand) return x;
381 auto shift=unsigned(-a.exponent);
382 auto integer=shift<64?a.significand>>shift:0;
383 bool guard=shift<=64 && ((a.significand>>(shift-1))&1);
384 bool sticky=shift>64 ? a.significand!=0 :
385 (shift>1 && (a.significand&((std::uint64_t{1}<<(shift-1))-1))!=0);
386 bool inexact=guard || sticky;
387 if(mode==rounding::to_odd && inexact) integer|=1;
388 else if((mode==rounding::nearest_even && guard && (sticky || (integer&1))) ||
389 (mode==rounding::downward && a.sign && inexact) ||
390 (mode==rounding::upward && !a.sign && inexact)) ++integer;
391 return round_pack<F>(a.sign,magnitude<1>{{integer}},0,mode,p);
392 }
393
394 template<class F>
395 constexpr bool equal_bits(typename F::bits_type x,typename F::bits_type y) noexcept {
396 return !is_nan<F>(x) && !is_nan<F>(y) && (x==y || (is_zero<F>(x) && is_zero<F>(y)));
397 }
398
399 template<class F>
400 constexpr bool less_bits(typename F::bits_type x,typename F::bits_type y) noexcept {
401 if(is_nan<F>(x) || is_nan<F>(y) || (is_zero<F>(x) && is_zero<F>(y))) return false;
402 bool sx=(x&F::sign_mask)!=0,sy=(y&F::sign_mask)!=0;
403 return sx!=sy?sx:sx?x>y:x<y;
404 }
405}
typename mask_traits< std::remove_cvref_t< T > >::type mask
Definition mask_traits.h:22