aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rhogenson@posteo.net>2022-08-08 19:43:22 -0700
committerRose Hogenson <rhogenson@posteo.net>2022-08-08 19:43:22 -0700
commit1317d29d0e4ac9cedd682bc43bca5e30a647244b (patch)
tree9b233a659de3fa0ace48971d97724c59e826e903
parentUnexport case-lambda. (diff)
downloadchromatopelma-1317d29d0e4ac9cedd682bc43bca5e30a647244b.tar.zst
Fix the bugs in the number library.
I did some extensive manual testing. Hopefully there are no bugs left :)
-rw-r--r--lib/scheme/base/80-number.csc388
1 files changed, 218 insertions, 170 deletions
diff --git a/lib/scheme/base/80-number.csc b/lib/scheme/base/80-number.csc
index 20f8db8..171bbd0 100644
--- a/lib/scheme/base/80-number.csc
+++ b/lib/scheme/base/80-number.csc
@@ -55,7 +55,7 @@
(define (small-int? obj)
- (call-builtin eq? 1 (call-builtin typeof obj)))
+ (call-builtin eq 1 (call-builtin typeof obj)))
(define (integer? obj)
@@ -66,15 +66,15 @@
(define (int=? n1 n2)
(cond
((and (small-int? n1) (small-int? n2))
- (call-builtin eq? n1 n2))
+ (call-builtin eq n1 n2))
((and (boxed-int? n1) (boxed-int? n2))
(let ((n1-digits (boxed-int-digits n1))
(n2-digits (boxed-int-digits n2)))
(and (boolean=? (boxed-int-positive? n1) (boxed-int-positive? n2))
- (call-builtin eq? (vector-length n1-digits) (vector-length n2-digits))
+ (call-builtin eq (vector-length n1-digits) (vector-length n2-digits))
(let loop ((i 0))
- (if (call-builtin int<? i (vector-length n1-digits))
- (and (call-builtin eq? (vector-ref n1-digits i) (vector-ref n2-digits i))
+ (if (call-builtin lt i (vector-length n1-digits))
+ (and (call-builtin eq (vector-ref n1-digits i) (vector-ref n2-digits i))
(loop (+ 1 i)))
#t)))))
(else #f)))
@@ -82,14 +82,14 @@
(define (int-positive? n)
(or (and (small-int? n)
- (call-builtin int<? 0 n))
+ (call-builtin lt 0 n))
(and (boxed-int? n)
(boxed-int-positive? n))))
(define (int-negative? n)
(or (and (small-int? n)
- (call-builtin int<? n 0))
+ (call-builtin lt n 0))
(and (boxed-int? n)
(not (boxed-int-positive? n)))))
@@ -97,23 +97,24 @@
(define (int<? n1 n2)
(cond
((and (small-int? n1) (small-int? n2))
- (call-builtin int<? n1 n2))
- ((and (int-negative? n1) (int-positive? n2)) #t)
- ((and (int-positive? n1) (int-negative? n2)) #f)
- ((and (int-negative? n1) (int-negative? n2)) (int<? (- n2) (- n1)))
- ((and (small-int? n1) (boxed-int? n2)) #t)
- ((and (boxed-int? n1) (small-int? n2)) #f)
+ (call-builtin lt n1 n2))
+ ((and (int-negative? n1) (int-negative? n2))
+ (int<? (- n2) (- n1)))
+ ((int-negative? n1) #t)
+ ((int-negative? n2) #f)
+ ((small-int? n1) #t)
+ ((small-int? n2) #f)
(else
(let* ((n1-digits (boxed-int-digits n1))
(n2-digits (boxed-int-digits n2))
(n1-digits-len (vector-length n1-digits))
(n2-digits-len (vector-length n2-digits)))
- (or (call-builtin int<? n1-digits-len n2-digits-len)
+ (or (call-builtin lt n1-digits-len n2-digits-len)
(and (int=? n1-digits-len n2-digits-len)
(let loop ((i (call-builtin sub n1-digits-len 1)))
(cond
- ((call-builtin int<? i 0) #f) ; n1 is equal to n2
- ((call-builtin int<? (vector-ref n1-digits i) (vector-ref n2-digits i)) #t)
+ ((call-builtin lt i 0) #f) ; n1 is equal to n2
+ ((call-builtin lt (vector-ref n1-digits i) (vector-ref n2-digits i)) #t)
((int=? (vector-ref n1-digits i) (vector-ref n2-digits i))
(loop (- i 1)))
(else #f))))))))) ; n1 > n2
@@ -134,7 +135,7 @@
(define (zero? obj)
- (call-builtin eq? 0 obj))
+ (call-builtin eq 0 obj))
(define (odd? n)
@@ -159,89 +160,84 @@
(define small-int-min (- #x4000000000000000)) ; -(2^62)
- (define digit-mask #x2000000000000000) ; 2^61. Digits are 61-bit unsigned numbers.
-
+ (define (digits n)
+ (if (small-int? n)
+ (vector n)
+ (boxed-int-digits n)))
- (define (split-carry-bit n)
- (values
- (call-builtin div n digit-mask)
- (call-builtin mod n digit-mask)))
+ ; adds two small positive ints.
+ (define (add2 x y)
+ (define z (call-builtin add x y))
+ (if (negative? z)
+ (values 1 (call-builtin add
+ 1
+ (call-builtin add
+ z
+ small-int-max)))
+ (values 0 z)))
- (define (denormalized-int n)
- (define-values (sgn n*) (if (int-positive? n)
- (values #t n)
- (values #f (- n))))
- (define-values (n2 n1) (split-carry-bit n*))
- (if (zero? n2)
- (make-boxed-int sgn (vector n1))
- (make-boxed-int sgn (vector n1 n2))))
-
- (define (big-int+ n1 n2)
+ (define (big-int+ x y)
(cond
- ((and (boxed-int-positive? n1) (boxed-int-negative? n2))
- (big-int- n1 (- n2)))
- ((and (boxed-int-negative? n1) (boxed-int-positive? n2))
- (big-int- n2 (- n1)))
- ((and (boxed-int-negative? n1) (boxed-int-negative? n2))
- (- (big-int+ (- n1) (- n2))))
+ ((and (int-negative? x) (int-negative? y))
+ (- (int+ (- x) (- y))))
+ ((int-negative? x)
+ (int- y (- x)))
+ ((int-negative? y)
+ (int- x (- y)))
(else
- (let* ((n1-digits (boxed-int-digits n1))
- (n2-digits (boxed-int-digits n2))
- (n1-ndigits (vector-length n1-digits))
- (n2-ndigits (vector-length n2-digits))
- (out-digits (make-vector (max n1-ndigits
- n2-ndigits))))
+ (let* ((x-digits (digits x))
+ (y-digits (digits y))
+ (x-ndigits (vector-length x-digits))
+ (y-ndigits (vector-length y-digits))
+ (out-digits (make-vector (max x-ndigits
+ y-ndigits))))
(let loop ((i 0)
(carry 0))
(cond
- ((and (int<? i n1-ndigits)
- (int<? i n2-ndigits))
- (let-values (((c d) (split-carry-bit
- (call-builtin add
- carry
- (call-builtin add
- (vector-ref n1-digits i)
- (vector-ref n2-digits i))))))
+ ((and (int<? i x-ndigits)
+ (int<? i y-ndigits))
+ (let*-values (((c1 d1) (add2
+ (vector-ref x-digits i)
+ (vector-ref y-digits i)))
+ ((c2 d2) (add2
+ d1
+ carry)))
+ (vector-set! out-digits i d2)
+ (loop (call-builtin add 1 i) (call-builtin add c1 c2))))
+ ((int<? i x-ndigits)
+ (let-values (((c d) (add2
+ carry
+ (vector-ref x-digits i))))
(vector-set! out-digits i d)
(loop (call-builtin add 1 i) c)))
- ((int<? i n1-ndigits)
- (let-values (((c d) (split-carry-bit
- (call-builtin add
- carry
- (vector-ref n1-digits i)))))
+ ((int<? i y-ndigits)
+ (let-values (((c d) (add2
+ carry
+ (vector-ref y-digits i))))
(vector-set! out-digits i d)
(loop (call-builtin add 1 i) c)))
- ((int<? i n2-ndigits)
- (let-values (((c d) (split-carry-bit
- (call-builtin add
- carry
- (vector-ref n2-digits i)))))))
((int-positive? carry)
(set! out-digits (vector-append out-digits (vector carry))))))
(make-boxed-int #t out-digits)))))
- (define (int+ n1 n2)
+ (define (int+ x y)
(cond
- ((and (small-int? n1)
- (small-int? n2))
- (let ((x (call-builtin add n1 n2))
- (n1-positive (int-positive? n1)))
- (if (and (boolean=? n1-positive (int-positive? n2))
- (not (boolean=? n1-positive (int-positive? x)))) ; overflow
- (big-int+ (denormalized-int n1) (denormalized-int n2))
- x)))
- (else
- (when (small-int? n1)
- (set! n1 (denormalized-int n1)))
- (when (small-int? n2)
- (set! n2 (denormalized-int n2)))
- (big-int+ n1 n2))))
+ ((and (small-int? x)
+ (small-int? y))
+ (let ((z (call-builtin add x y))
+ (x-positive (int-positive? x)))
+ (if (and (boolean=? x-positive (int-positive? y))
+ (not (boolean=? x-positive (int-positive? z)))) ; overflow
+ (big-int+ x y)
+ z)))
+ (else (big-int+ x y))))
(define (remove-leading-zeros n)
+ (define digits (boxed-int-digits n))
(define leading-zeros
(let loop ((i (int- (vector-length digits) 1))
(n 0))
@@ -253,17 +249,15 @@
(define (normalize n)
+ (set! n (remove-leading-zeros n))
(define digits (boxed-int-digits n))
- (if (and (int<? (vector-length digits) 3)
- (int<? (vector-ref digits 1) (if (int-positive? n)
- 2
- 3))) ; int min requires special handling
- (if (int-negative? n)
- (- (+ (vector-ref digits 0)
- (* digit-mask (vector-ref digits 1))))
- (+ (vector-ref digits 0)
- (* digit-mask (vector-ref digits 1))))
- n)) ; I have no doubt that there is a bug in this procedure.
+ (define len (vector-length digits))
+ (cond
+ ((call-builtin eq 0 len)
+ 0)
+ ((call-builtin eq 1 len)
+ (vector-ref digits 0))
+ (else n)))
(define (big-int- n1 n2)
@@ -277,12 +271,11 @@
((int<? n1 n2)
(- (big-int- n2 n1)))
(else
- (let* ((n1-digits (boxed-int-digits n1))
- (n2-digits (boxed-int-digits n2))
+ (let* ((n1-digits (digits n1))
+ (n2-digits (digits n2))
(n1-ndigits (vector-length n1-digits))
(n2-ndigits (vector-length n2-digits))
- (out-digits (make-vector (max n1-ndigits
- n2-ndigits))))
+ (out-digits (make-vector n1-ndigits))) ; n.b.: n1 is bigger
(let loop ((i 0)
(carry 0))
(cond
@@ -295,12 +288,30 @@
(vector-ref n2-digits i)))))
(if (int-negative? x)
(begin
- (vector-set! out-digits i (call-builtin add x digit-mask))
+ (vector-set! out-digits i
+ (call-builtin add
+ 1
+ (call-builtin add x small-int-max)))
+ (loop (+ 1 i) -1))
+ (begin
+ (vector-set! out-digits i x)
+ (loop (+ 1 i) 0)))))
+ ((int<? i n1-ndigits)
+ (let ((x (call-builtin add
+ carry
+ (vector-ref n1-digits i))))
+ (if (int-negative? x)
+ (begin
+ (vector-set! out-digits i
+ (call-builtin add
+ 1
+ (call-builtin add x small-int-max)))
(loop (+ 1 i) -1))
(begin
(vector-set! out-digits i x)
(loop (+ 1 i) 0)))))))
- (normalize (remove-leading-zeros (make-boxed-int #t out-digits)))))))
+
+ (normalize (make-boxed-int #t out-digits))))))
(define (int- n1 n2)
@@ -311,30 +322,61 @@
(if (or (boolean=? n1-positive (int-positive? n2))
(boolean=? n1-positive (int-positive? x)))
x
- (big-int- (denormalized-int n1) (denormalized-int n2)))))
+ (big-int- n1 n2))))
(else
- (when (small-int? n1)
- (set! n1 (denormalized-int n1)))
- (when (small-int? n2)
- (set! n2 (denormalized-int n2)))
(big-int- n1 n2))))
+ (define half-word-mask #x80000000)
+
+
+ ; multiplies two small positive ints.
+ (define (mul2 x y)
+ ; Split each argument into half-words.
+ (define x0 (call-builtin mod x half-word-mask))
+ (define x1 (call-builtin div x half-word-mask))
+ (define y0 (call-builtin mod y half-word-mask))
+ (define y1 (call-builtin div y half-word-mask))
+
+ ; Do the grade-school multiplication algorithm.
+ (define z0 (call-builtin mul x0 y0))
+ (define-values (z1-hi z1-lo)
+ (add2 (call-builtin mul x1 y0)
+ (call-builtin mul x0 y1)))
+ (define z2 (call-builtin mul x1 y1))
+
+ ; Shift z1 left by a half-word.
+ (define z1*-lo (call-builtin mul
+ half-word-mask
+ (call-builtin mod z1-lo half-word-mask)))
+ (define z1*-hi (call-builtin add
+ (call-builtin mul z1-hi half-word-mask)
+ (call-builtin div z1-lo half-word-mask)))
+
+ ; Add z0 + z1* + z2*word-size
+ (define-values (carry result-lo) (add2 z0 z1*-lo))
+ (values
+ (call-builtin add
+ carry
+ (call-builtin add
+ z1*-hi
+ z2))
+ result-lo))
+
+
(define (split-in-half n m)
- (define digits (boxed-int-digits n))
- (if (int<? m (vector-length digits))
+ (define n-digits (digits n))
+ (if (int<? m (vector-length n-digits))
(values
- (normalize (make-boxed-int #t (vector-copy m)))
- (normalize (make-boxed-int #t (vector-copy 0 m))))
+ (normalize (make-boxed-int #t (vector-copy n-digits m)))
+ (normalize (make-boxed-int #t (vector-copy n-digits 0 m))))
(values
0
n)))
(define (lshift n m)
- (when (small-int? n)
- (set! n (denormalized-int n)))
- (let* ((old-digits (boxed-int-digits m))
+ (let* ((old-digits (digits n))
(new-digits (make-vector (int+ m (vector-length old-digits)))))
(vector-fill! new-digits 0 0 m)
(vector-copy! new-digits m old-digits)
@@ -343,28 +385,37 @@
; This is Karatsuba's algorithm.
(define (big-int* x y)
+ (define m (call-builtin div
+ (max (vector-length (digits x)) (vector-length (digits y)))
+ 2))
+ (define-values (x1 x0) (split-in-half x m))
+ (define-values (y1 y0) (split-in-half y m))
+ (define z0 (* x0 y0))
+ (define z2 (* x1 y1))
+ (define z1 (- (* (+ x1 x0)
+ (+ y1 y0))
+ z2
+ z0))
+ (define ans
+ (+ (lshift z2 (call-builtin mul 2 m))
+ (lshift z1 m)
+ z0))
+ ans)
+
+
+ (define (slow-int* x y)
(cond
- ((and (boxed-int-positive? x) (boxed-int-positive? y))
- (define m (call-builtin div
- (max (vector-length (boxed-int-digits x)) (vector-length (boxed-int-digits y)))
- 2))
- (define-values (x1 x0) (split-in-half x m))
- (define-values (y1 y0) (split-in-half y m))
- (define z0 (* x0 y0))
- (define z2 (* x1 y1))
- (define z1 (- (* (+ x1 x0)
- (+ y1 y0))
- z2
- z0))
- (+ (lshift z2 (call-builtin mul 2 m))
- (lshift z1 m)
- z0))
- ((and (boxed-int-negative? x) (boxed-int-negative? y))
- (big-int* (- x) (- y)))
- ((boxed-int-negative? x)
- (- (big-int* (- x) y)))
- (else ; y is negative
- (- (big-int* x (- y))))))
+ ((and (int-negative? x) (int-negative? y))
+ (int* (- x) (- y)))
+ ((int-negative? x)
+ (- (int* (- x) y)))
+ ((int-negative? y)
+ (- (int* x (- y))))
+ ((and (small-int? x) (small-int? y))
+ (let-values (((z1 z0) (mul2 x y)))
+ (make-boxed-int #t (vector z0 z1))))
+ (else
+ (big-int* x y))))
(define (int* n1 n2)
@@ -374,14 +425,10 @@
(if (and (not (zero? n1))
(not (int=? (call-builtin div x n1)
n2))) ; overflow
- (big-int* (denormalized-int n1) (denormalized-int n2))
+ (slow-int* n1 n2)
x)))
(else
- (when (small-int? n1)
- (set! n1 (denormalized-int n1)))
- (when (small-int? n2)
- (set! n2 (denormalized-int n2)))
- (big-int* n1 n2))))
+ (slow-int* n1 n2))))
(define (left-index v i)
@@ -390,38 +437,41 @@
(define (find-beta d m)
(let loop ((lo 0)
- (hi digit-mask))
- (define guess (call-builtin div (- hi lo) 2))
+ (hi small-int-max))
+ (define guess (+ lo (call-builtin div (- hi lo) 2)))
(define check (- d (* m guess)))
(cond
+ ((int>? lo hi)
+ (error "binary search is hard"))
((negative? check)
; guess was too big
- (loop lo guess))
+ (loop lo (- guess 1)))
((int<? check m)
; got it!
guess)
- (else
+ ((int<=? lo hi)
; guess was too small
- (loop (+ guess 1) hi)))))
+ (loop (+ guess 1) hi))
+ (else (error "binary search is broken" d m)))))
(define (big-int/ n m)
- (define n-digits (boxed-int-digits n))
- (define m-digits (boxed-int-digits m))
+ (define n-digits (digits n))
+ (define m-digits (digits m))
(define k (vector-length n-digits))
(define l (vector-length m-digits))
(if (int<? k l)
- (values 0 (normalize n))
+ (values 0 n)
(let loop ((i (- l 1))
(q 0)
(r (normalize
- (make-boxed-int #t (vector-copy (boxed-int-digits n) (- k l 1)))))) ; last l-1 digits of n
+ (make-boxed-int #t (vector-copy n-digits (- k (- l 1))))))) ; last l-1 digits of n
(if (int<? i k)
- (let* ((d (+ (* digit-mask r)
- (left-index n-digits i)))
+ (let* ((d (+ (lshift r 1)
+ (left-index n-digits i)))
(beta (find-beta d m)))
(loop (+ 1 i)
- (+ (* digit-mask q)
+ (+ (lshift q 1)
beta)
(- d (* m beta))))
(values q r)))))
@@ -430,7 +480,7 @@
(define (truncate/ n1 n2)
(cond
((and (small-int? n1) (small-int? n2))
- (if (= (-1 n2))
+ (if (= -1 n2)
; To avoid dividing INT_MIN by -1,
; we just convert every division by -1 into a negation.
(values (- n1) 0)
@@ -447,11 +497,6 @@
(let-values (((q r) (truncate/ n1 (- n2))))
(values (- q) r)))
(else
- (when (small-int? n1)
- (set! n1 (denormalized-int n1)))
- (when (small-int? n2)
- (set! n2 (denormalized-int n2)))
- ; big and positive
(big-int/ n1 n2))))
@@ -576,11 +621,11 @@
(int<? x1 x2)
(let* ((d1 (denominator x1))
(d2 (denominator x2))
- (g (gcd d1 d2)))
+ (> (gcd d1 d2)))
(int<? (* (numerator x1)
- (quotient d2 g))
+ (quotient d2 >))
(* (numerator x2)
- (quotient d1 g))))))
+ (quotient d1 >))))))
(define <
@@ -640,7 +685,7 @@
(define (min x1 . xs)
(let loop ((xs xs)
- (m m1))
+ (m x1))
(if (null? xs)
m
(loop (cdr xs)
@@ -658,9 +703,9 @@
(else
(let* ((d1 (denominator x1))
(d2 (denominator x2))
- (g (gcd d1 d2))
- (s1 (quotient d2 g))
- (s2 (quotient d1 g)))
+ (> (gcd d1 d2))
+ (s1 (quotient d2 >))
+ (s2 (quotient d1 >)))
(/ (+ (* s1 (numerator x1))
(* s2 (numerator x2)))
(* d1 s1))))))
@@ -695,6 +740,8 @@
(case-lambda
((z)
(cond
+ ((int=? z small-int-min)
+ #x4000000000000000)
((small-int? z)
(call-builtin sub 0 z))
((boxed-int? z)
@@ -717,14 +764,14 @@
(define (normalize-quotient q)
- (when (negative? d)
+ (when (negative? q)
(set! q (make-quotient (- (numerator q))
(- (denominator q)))))
(let* ((n (numerator q))
(d (denominator q))
- (g (gcd n d)))
- (make-quotient (quotient n g)
- (quotient d g))))
+ (> (gcd n d)))
+ (make-quotient (quotient n >)
+ (quotient d >))))
(define /
@@ -738,8 +785,10 @@
(set! z1 (- z1))
(set! z2 (- z2)))
(let ((g (gcd z1 z2)))
- (make-quotient (quotient z1 g)
- (quotient z2 g))))
+ (if (= g z2)
+ (quotient z1 g)
+ (make-quotient (quotient z1 g)
+ (quotient z2 g)))))
(else
(* z1
(/ (denominator z2)
@@ -801,12 +850,11 @@
(define (exact-integer-sqrt k)
(if (zero? k)
(values 0 0)
- (let loop ((x (/ k 4))) ; initial estimate
- (define y (/ (+ x (/ k x))
- 2)) ; 1/2 (x + k/x)
- (if (> 1 (abs (- x y)))
- (let ((s (truncate y)))
- (values s (- k (* s s))))
+ (let loop ((x (quotient k 2))) ; initial estimate
+ (define y (quotient (+ x (quotient k x))
+ 2))
+ (if (>= y x)
+ (values x (- k (square x)))
(loop y)))))