1use binius_frontend::{CircuitBuilder, Wire};
3
4use super::utils::*;
5
6pub fn fp64_mul_prepare(
21 b: &CircuitBuilder,
22 pa: &Fp64Parts,
23 pb: &Fp64Parts,
24) -> (Wire, Wire, Wire, Wire) {
25 let bias = b.add_constant_64(1023);
26
27 let (m_a, exp_eff_a) = fp64_sig53_and_exp(b, pa);
29 let (m_b, exp_eff_b) = fp64_sig53_and_exp(b, pb);
30
31 let exp_sum = b.iadd(exp_eff_a, exp_eff_b).0;
33 let exp_pre = isub(b, exp_sum, bias);
34
35 let sign = b.bxor(pa.sign, pb.sign);
37
38 (m_a, m_b, exp_pre, sign)
39}
40
41pub fn fp64_mul_make_round_base(b: &CircuitBuilder, m_a: Wire, m_b: Wire) -> (Wire, Wire) {
61 let (hi, lo) = b.imul(m_a, m_b); let top105_bit01 = bit_lsb(b, hi, 41); let top105_sel = bit_msb01(b, hi, 41); let y41 = shr128_to_u64_const(b, hi, lo, 41);
70 let sticky41 = sticky_from_low_k(b, lo, 41);
71 let y42 = shr128_to_u64_const(b, hi, lo, 42);
73 let sticky42 = sticky_from_low_k(b, lo, 42);
74
75 let y = b.select(top105_sel, y42, y41);
77 let sticky01 = b.select(top105_sel, sticky42, sticky41);
78
79 let sig_round_base = b.bor(y, sticky01);
83
84 (sig_round_base, top105_bit01) }
86
87pub fn fp64_mul_finish_specials(
103 b: &CircuitBuilder,
104 pa: &Fp64Parts,
105 pb: &Fp64Parts,
106 sign_xor: Wire,
107 finite_result: Wire,
108) -> Wire {
109 let qnan = b.add_constant_64(0x7FF8_0000_0000_0000);
110 let exp_2047 = b.add_constant_64(0x7FF);
111 let inf_payload = b.shl(exp_2047, 52);
112
113 let any_nan = b.bor(pa.is_nan, pb.is_nan);
114 let any_inf = b.bor(pa.is_inf, pb.is_inf);
115 let any_zero = b.bor(pa.is_zero, pb.is_zero);
116 let inf_times_zero = b.bor(b.band(pa.is_inf, pb.is_zero), b.band(pb.is_inf, pa.is_zero));
117
118 let nan_mask = b.bor(any_nan, inf_times_zero);
119
120 let sign_msb = b.band(sign_xor, b.add_constant_64(1u64 << 63));
121 let packed_inf = b.bor(sign_msb, inf_payload);
122
123 let with_zero = b.select(any_zero, sign_msb, finite_result);
125 let with_inf = b.select(any_inf, packed_inf, with_zero);
126 b.select(nan_mask, qnan, with_inf)
127}
128
129pub fn float64_mul(builder: &CircuitBuilder, a: Wire, b: Wire) -> Wire {
143 let pa = fp64_unpack(builder, a);
145 let pb = fp64_unpack(builder, b);
146
147 let (m_a, m_b, exp_pre, sign_msb) = fp64_mul_prepare(builder, &pa, &pb);
149
150 let (sig_round_base_uncut, norm_shift1_bit01) = fp64_mul_make_round_base(builder, m_a, m_b);
152
153 let exp_round_base = builder.iadd(exp_pre, norm_shift1_bit01).0;
155
156 let (sig_round_base, exp_for_round, exp_lt_1) =
158 fp64_underflow_shift(builder, sig_round_base_uncut, exp_round_base);
159
160 let (mant_final_53, exp_after_round, mant_overflow_mask) =
162 fp64_round_rne(builder, sig_round_base, exp_for_round);
163
164 let stayed_sub_mask = builder.band(exp_lt_1, builder.bnot(mant_overflow_mask));
166 let finite_or_inf =
167 fp64_pack_finite_or_inf(builder, sign_msb, mant_final_53, exp_after_round, stayed_sub_mask);
168
169 fp64_mul_finish_specials(builder, &pa, &pb, sign_msb, finite_or_inf)
171}
172
173#[cfg(test)]
174mod tests {
175 use binius_core::word::Word;
176
177 use super::*;
178 use crate::float64::utils::tests::{
179 f64_bits_semantic_eq, ref_fp64_pack_finite_or_inf, ref_fp64_round_rne,
180 ref_fp64_underflow_shift, ref_fp64_unpack,
181 };
182
183 fn ref_fp64_sig53_and_exp(exp: u64, frac: u64, is_norm: bool) -> (u64, u64) {
184 if is_norm {
185 ((1u64 << 52) | frac, exp)
186 } else {
187 (frac, 1)
188 }
189 }
190
191 #[allow(clippy::too_many_arguments)]
192 fn ref_fp64_mul_prepare(
193 pa_sign: u64,
194 pa_exp: u64,
195 pa_frac: u64,
196 pa_is_norm: bool,
197 pb_sign: u64,
198 pb_exp: u64,
199 pb_frac: u64,
200 pb_is_norm: bool,
201 ) -> (u64, u64, u64, u64) {
202 let bias = 1023u64;
203
204 let (m_a, exp_eff_a) = ref_fp64_sig53_and_exp(pa_exp, pa_frac, pa_is_norm);
205 let (m_b, exp_eff_b) = ref_fp64_sig53_and_exp(pb_exp, pb_frac, pb_is_norm);
206
207 let exp_sum = exp_eff_a.wrapping_add(exp_eff_b);
208 let exp_pre = exp_sum.wrapping_sub(bias);
209 let sign = pa_sign ^ pb_sign;
210
211 (m_a, m_b, exp_pre, sign)
212 }
213
214 fn ref_shr128_to_u64_const(hi: u64, lo: u64, s: u32) -> u64 {
215 debug_assert!(s > 0 && s < 64);
216 let p = ((hi as u128) << 64) | (lo as u128);
217 (p >> s) as u64
218 }
219
220 fn ref_sticky_from_low_k(lo: u64, k: u32) -> u64 {
221 debug_assert!(k < 64);
222 let mask = (1u64 << k) - 1;
223 if (lo & mask) != 0 { 1 } else { 0 }
224 }
225
226 fn ref_fp64_mul_make_round_base(m_a: u64, m_b: u64) -> (u64, u64) {
227 let p = (m_a as u128) * (m_b as u128);
228 let hi = (p >> 64) as u64;
229 let lo = p as u64;
230
231 let top105_bit01 = (hi >> 41) & 1;
233
234 let (y, sticky01) = if top105_bit01 == 1 {
235 let y = ref_shr128_to_u64_const(hi, lo, 42);
237 let sticky = ref_sticky_from_low_k(lo, 42);
238 (y, sticky)
239 } else {
240 let y = ref_shr128_to_u64_const(hi, lo, 41);
242 let sticky = ref_sticky_from_low_k(lo, 41);
243 (y, sticky)
244 };
245
246 let new_b0 = (y & 1) | sticky01;
248 let sig_round_base = (y & !1) | new_b0;
249
250 (sig_round_base, top105_bit01)
251 }
252
253 #[allow(clippy::too_many_arguments)]
254 fn ref_fp64_mul_finish_specials(
255 pa_is_nan: u64,
256 pa_is_inf: u64,
257 pa_is_zero: u64,
258 pb_is_nan: u64,
259 pb_is_inf: u64,
260 pb_is_zero: u64,
261 sign_xor: u64,
262 finite_result: u64,
263 ) -> u64 {
264 let qnan = 0x7FF8_0000_0000_0000u64;
265 let inf_payload = 0x7FFu64 << 52;
266
267 let any_nan = (pa_is_nan | pb_is_nan) != 0;
268 let any_inf = (pa_is_inf | pb_is_inf) != 0;
269 let any_zero = (pa_is_zero | pb_is_zero) != 0;
270 let inf_times_zero =
271 ((pa_is_inf != 0) && (pb_is_zero != 0)) || ((pb_is_inf != 0) && (pa_is_zero != 0));
272
273 let nan_case = any_nan || inf_times_zero;
274
275 if nan_case {
276 return qnan;
277 }
278 if any_inf {
279 return sign_xor | inf_payload;
280 }
281 if any_zero {
282 return sign_xor; }
284 finite_result
285 }
286
287 fn ref_float64_mul(a_bits: u64, b_bits: u64) -> u64 {
288 let pa = ref_fp64_unpack(a_bits);
290 let pb = ref_fp64_unpack(b_bits);
291
292 let pa_is_zero = (pa.exp == 0) && (pa.frac == 0);
295 let pb_is_zero = (pb.exp == 0) && (pb.frac == 0);
296
297 let any_nan = (pa.is_nan | pb.is_nan) != 0;
299 let inf_times_zero = ((pa.is_inf != 0) && pb_is_zero) || ((pb.is_inf != 0) && pa_is_zero);
300 if any_nan || inf_times_zero {
301 return 0x7FF8_0000_0000_0000u64; }
303
304 let any_inf = (pa.is_inf | pb.is_inf) != 0;
306 if any_inf {
307 let sign_xor_msb = pa.sign ^ pb.sign; return sign_xor_msb | (0x7FFu64 << 52);
309 }
310
311 let any_zero = pa_is_zero || pb_is_zero;
313 if any_zero {
314 let sign_xor_msb = pa.sign ^ pb.sign; return sign_xor_msb; }
317
318 let sign_xor = pa.sign ^ pb.sign; let (m_a, exp_eff_a) = if pa.is_norm != 0 {
323 ((1u64 << 52) | pa.frac, pa.exp)
324 } else {
325 (pa.frac, 1)
326 };
327 let (m_b, exp_eff_b) = if pb.is_norm != 0 {
328 ((1u64 << 52) | pb.frac, pb.exp)
329 } else {
330 (pb.frac, 1)
331 };
332
333 let exp_pre = exp_eff_a.wrapping_add(exp_eff_b).wrapping_sub(1023);
335
336 let p = (m_a as u128) * (m_b as u128);
338 let hi = (p >> 64) as u64;
339 let lo = p as u64;
340
341 let top105_bit = (hi >> 41) & 1;
343 let (sig_round_base_uncut, norm_shift) = if top105_bit == 1 {
344 let y = ((hi as u128) << 64 | lo as u128) >> 42;
346 let sticky = if (lo & ((1u64 << 42) - 1)) != 0 { 1 } else { 0 };
347 let y64 = y as u64;
348 let new_b0 = (y64 & 1) | sticky;
349 ((y64 & !1) | new_b0, 1)
350 } else {
351 let y = ((hi as u128) << 64 | lo as u128) >> 41;
353 let sticky = if (lo & ((1u64 << 41) - 1)) != 0 { 1 } else { 0 };
354 let y64 = y as u64;
355 let new_b0 = (y64 & 1) | sticky;
356 ((y64 & !1) | new_b0, 0)
357 };
358
359 let exp_round_base = exp_pre.wrapping_add(norm_shift);
361
362 let (sig_round_base, exp_for_round, exp_lt_1) =
364 ref_fp64_underflow_shift(sig_round_base_uncut, exp_round_base);
365
366 let (mant_final_53, exp_after_round, mant_overflow_mask) =
368 ref_fp64_round_rne(sig_round_base, exp_for_round);
369
370 let stayed_sub_mask = if exp_lt_1 != 0 && mant_overflow_mask == 0 {
372 1u64 << 63
373 } else {
374 0
375 };
376 ref_fp64_pack_finite_or_inf(sign_xor, mant_final_53, exp_after_round, stayed_sub_mask)
377 }
378
379 #[test]
380 fn test_fp64_mul_prepare() {
381 let test_cases = [
382 (1.0f64.to_bits(), 2.0f64.to_bits()),
384 (1.5f64.to_bits(), 2.5f64.to_bits()),
385 ((-1.0f64).to_bits(), 2.0f64.to_bits()),
386 (1.0f64.to_bits(), (-3.0f64).to_bits()),
387 (f64::MIN_POSITIVE.to_bits(), 2.0f64.to_bits()),
388 ];
389
390 for (a_bits, b_bits) in test_cases {
391 let builder = CircuitBuilder::new();
392 let a_wire = builder.add_inout();
393 let b_wire = builder.add_inout();
394 let pa = fp64_unpack(&builder, a_wire);
395 let pb = fp64_unpack(&builder, b_wire);
396 let (m_a, m_b, exp_pre, sign) = fp64_mul_prepare(&builder, &pa, &pb);
397
398 let expected_m_a = builder.add_inout();
399 let expected_m_b = builder.add_inout();
400 let expected_exp_pre = builder.add_inout();
401 let expected_sign = builder.add_inout();
402 let mask = builder.add_constant_64(1u64 << 63);
403
404 builder.assert_eq("m_a", m_a, expected_m_a);
405 builder.assert_eq("m_b", m_b, expected_m_b);
406 builder.assert_eq("exp_pre", exp_pre, expected_exp_pre);
407 builder.assert_eq("sign", builder.band(sign, mask), expected_sign);
408
409 let circuit = builder.build();
410 let mut w = circuit.new_witness_filler();
411 w[a_wire] = Word(a_bits);
412 w[b_wire] = Word(b_bits);
413
414 let pa_sign = a_bits;
416 let pa_exp = (a_bits >> 52) & 0x7FF;
417 let pa_frac = a_bits & ((1u64 << 52) - 1);
418 let pa_is_norm = pa_exp != 0 && pa_exp != 0x7FF;
419
420 let pb_sign = b_bits;
421 let pb_exp = (b_bits >> 52) & 0x7FF;
422 let pb_frac = b_bits & ((1u64 << 52) - 1);
423 let pb_is_norm = pb_exp != 0 && pb_exp != 0x7FF;
424
425 let (ref_m_a, ref_m_b, ref_exp_pre, ref_sign) = ref_fp64_mul_prepare(
426 pa_sign, pa_exp, pa_frac, pa_is_norm, pb_sign, pb_exp, pb_frac, pb_is_norm,
427 );
428
429 w[expected_m_a] = Word(ref_m_a);
430 w[expected_m_b] = Word(ref_m_b);
431 w[expected_exp_pre] = Word(ref_exp_pre);
432 w[expected_sign] = Word(ref_sign & (1u64 << 63));
433
434 circuit.populate_wire_witness(&mut w).unwrap();
435 let cs = circuit.constraint_system();
436 cs.verify(&w.into_value_vec()).unwrap();
437 }
438 }
439
440 #[test]
441 fn test_fp64_mul_make_round_base() {
442 let test_cases = [
443 (1u64 << 52, 1u64 << 52), ((1u64 << 52) + (1u64 << 51), 1u64 << 52), ((1u64 << 52) + (1u64 << 51), (1u64 << 52) + (1u64 << 51)), (((1u64 << 53) - 1), ((1u64 << 53) - 1)), ];
449
450 for (m_a_val, m_b_val) in test_cases {
451 let builder = CircuitBuilder::new();
452 let m_a = builder.add_inout();
453 let m_b = builder.add_inout();
454 let (sig_round_base, norm_shift1) = fp64_mul_make_round_base(&builder, m_a, m_b);
455
456 let expected_sig = builder.add_inout();
457 let expected_norm = builder.add_inout();
458
459 builder.assert_eq("sig_round_base", sig_round_base, expected_sig);
460 builder.assert_eq("norm_shift1", norm_shift1, expected_norm);
461
462 let circuit = builder.build();
463 let mut w = circuit.new_witness_filler();
464 w[m_a] = Word(m_a_val);
465 w[m_b] = Word(m_b_val);
466
467 let (ref_sig, ref_norm) = ref_fp64_mul_make_round_base(m_a_val, m_b_val);
468 w[expected_sig] = Word(ref_sig);
469 w[expected_norm] = Word(ref_norm);
470
471 circuit.populate_wire_witness(&mut w).unwrap();
472 let cs = circuit.constraint_system();
473 cs.verify(&w.into_value_vec()).unwrap();
474 }
475 }
476
477 #[test]
478 fn test_fp64_mul_finish_specials() {
479 let test_cases = [
480 (0, 0, 0, 0, 0, 0, 0, 0x4000000000000000u64), (1u64 << 63, 0, 0, 0, 0, 0, 0, 0x4000000000000000u64), (0, 0, 0, 1u64 << 63, 0, 0, 1u64 << 63, 0x4000000000000000u64), (0, 1u64 << 63, 0, 0, 0, 1u64 << 63, 0, 0x4000000000000000u64), (0, 0, 1u64 << 63, 0, 1u64 << 63, 0, 1u64 << 63, 0x4000000000000000u64), (0, 1u64 << 63, 0, 0, 0, 0, 0, 0x4000000000000000u64), (0, 1u64 << 63, 0, 0, 0, 0, 1u64 << 63, 0x4000000000000000u64), (0, 0, 1u64 << 63, 0, 0, 0, 0, 0x4000000000000000u64), (0, 0, 1u64 << 63, 0, 0, 0, 1u64 << 63, 0x4000000000000000u64), ];
494
495 for (
496 pa_is_nan,
497 pa_is_inf,
498 pa_is_zero,
499 pb_is_nan,
500 pb_is_inf,
501 pb_is_zero,
502 sign_xor,
503 finite_result,
504 ) in test_cases
505 {
506 let builder = CircuitBuilder::new();
507 let pa = Fp64Parts {
508 sign: builder.add_inout(),
509 exp: builder.add_inout(),
510 frac: builder.add_inout(),
511 is_nan: builder.add_inout(),
512 is_inf: builder.add_inout(),
513 is_zero: builder.add_inout(),
514 is_sub: builder.add_inout(),
515 is_norm: builder.add_inout(),
516 };
517 let pb = Fp64Parts {
518 sign: builder.add_inout(),
519 exp: builder.add_inout(),
520 frac: builder.add_inout(),
521 is_nan: builder.add_inout(),
522 is_inf: builder.add_inout(),
523 is_zero: builder.add_inout(),
524 is_sub: builder.add_inout(),
525 is_norm: builder.add_inout(),
526 };
527 let sign_xor_wire = builder.add_inout();
528 let finite_result_wire = builder.add_inout();
529
530 let result =
531 fp64_mul_finish_specials(&builder, &pa, &pb, sign_xor_wire, finite_result_wire);
532 let expected_result = builder.add_inout();
533 builder.assert_eq("finish_result", result, expected_result);
534
535 let circuit = builder.build();
536 let mut w = circuit.new_witness_filler();
537
538 w[pa.sign] = Word(0);
540 w[pa.exp] = Word(0);
541 w[pa.frac] = Word(0);
542 w[pa.is_nan] = Word(pa_is_nan);
543 w[pa.is_inf] = Word(pa_is_inf);
544 w[pa.is_zero] = Word(pa_is_zero);
545 w[pa.is_sub] = Word(0);
546 w[pa.is_norm] = Word(0);
547
548 w[pb.sign] = Word(0);
549 w[pb.exp] = Word(0);
550 w[pb.frac] = Word(0);
551 w[pb.is_nan] = Word(pb_is_nan);
552 w[pb.is_inf] = Word(pb_is_inf);
553 w[pb.is_zero] = Word(pb_is_zero);
554 w[pb.is_sub] = Word(0);
555 w[pb.is_norm] = Word(0);
556
557 w[sign_xor_wire] = Word(sign_xor);
558 w[finite_result_wire] = Word(finite_result);
559
560 let ref_result = ref_fp64_mul_finish_specials(
561 pa_is_nan,
562 pa_is_inf,
563 pa_is_zero,
564 pb_is_nan,
565 pb_is_inf,
566 pb_is_zero,
567 sign_xor,
568 finite_result,
569 );
570 w[expected_result] = Word(ref_result);
571
572 circuit.populate_wire_witness(&mut w).unwrap();
573 let cs = circuit.constraint_system();
574 cs.verify(&w.into_value_vec()).unwrap();
575 }
576 }
577
578 #[test]
579 fn test_float64_mul() {
580 let small_values = [
581 f64::MIN_POSITIVE, f64::MIN_POSITIVE * 2.0, 5e-324, 1e-308, 2.2250738585072014e-308, 1e-200, ];
588
589 let medium_values = [
590 1.0,
591 2.0,
592 1.5,
593 2.5,
594 std::f64::consts::PI,
595 10.0,
596 1000.0,
597 0.1,
598 0.25,
599 0.333333333,
600 ];
601
602 let large_values = [
603 1e100,
604 1e200,
605 1e307, f64::MAX / 2.0, 1.7976931348623155e308, 1e50, ];
610
611 let mut test_cases = Vec::new();
612
613 test_cases.extend([
615 (1.0, 2.0),
616 (1.5, 2.0),
617 (-1.0, 2.0),
618 (1.0, -2.0),
619 (-1.5, -2.5),
620 ]);
621
622 test_cases.extend([(0.0, 1.0), (-0.0, 1.0), (0.0, -1.0), (-0.0, -1.0)]);
624 for &val in &[small_values[0], medium_values[0], large_values[0]] {
625 test_cases.extend([
626 (val, 0.0),
627 (-val, 0.0),
628 (val, f64::INFINITY),
629 (val, f64::NAN),
630 ]);
631 }
632
633 test_cases.extend([
635 (f64::INFINITY, 1.0),
636 (-f64::INFINITY, 1.0),
637 (f64::INFINITY, -1.0),
638 (f64::INFINITY, 0.0), (0.0, f64::INFINITY), ]);
641
642 test_cases.extend([(f64::NAN, 1.0), (1.0, f64::NAN)]);
644
645 for &a in &small_values {
647 for &b in &small_values[..3] {
648 let native_product = a * b;
652 if native_product == 0.0 && (a != 0.0 && b != 0.0) {
653 continue; }
655 test_cases.push((a, b));
656 test_cases.push((-a, b));
657 test_cases.push((a, -b));
658 }
659 }
660
661 for &a in &medium_values {
663 for &b in &medium_values[..4] {
664 test_cases.push((a, b));
666 test_cases.push((-a, b));
667 }
668 }
669
670 for &a in &large_values[..3] {
672 for &b in &[1.0, 0.1, 2.0] {
674 test_cases.push((a, b));
676 test_cases.push((-a, b));
677 }
678 }
679
680 for &small in &small_values[..2] {
683 for &medium in &medium_values[..3] {
684 test_cases.push((small, medium));
685 test_cases.push((-small, medium));
686 }
687 }
688
689 for &medium in &[1.0, 2.0, 0.5] {
691 for &large in &large_values[..2] {
692 test_cases.push((medium, large));
693 test_cases.push((-medium, large));
694 }
695 }
696
697 for &small in &small_values[..2] {
699 for &large in &large_values[..2] {
700 test_cases.push((small, large));
701 test_cases.push((-small, -large));
702 }
703 }
704
705 for (i, (a_val, b_val)) in test_cases.iter().copied().enumerate() {
707 let builder = CircuitBuilder::new();
708 let a = builder.add_inout();
709 let b = builder.add_inout();
710 let result = float64_mul(&builder, a, b);
711 let expected = builder.add_inout();
712 builder.assert_eq(format!("float64_mul_case_{}", i), result, expected);
713
714 let circuit = builder.build();
715 let mut w = circuit.new_witness_filler();
716 w[a] = Word(a_val.to_bits());
717 w[b] = Word(b_val.to_bits());
718
719 let ref_result = ref_float64_mul(a_val.to_bits(), b_val.to_bits());
721 w[expected] = Word(ref_result);
722
723 circuit.populate_wire_witness(&mut w).unwrap();
724 let cs = circuit.constraint_system();
725
726 let circuit_result = w[result].0;
728 let result_constraints = cs.verify(&w.into_value_vec());
729
730 if let Err(e) = result_constraints {
732 panic!("Constraint verification failed for case {}: {:?}", i, e);
733 }
734
735 let native_result = (a_val * b_val).to_bits();
737 if !f64_bits_semantic_eq(circuit_result, native_result) {
738 panic!(
739 "Case {} ({:e} * {:e}): Circuit result 0x{:016x} doesn't match native result 0x{:016x} semantically",
740 i, a_val, b_val, circuit_result, native_result
741 );
742 }
743 }
744 }
745}