diff options
Diffstat (limited to 'lib')
| -rw-r--r-- | lib/scheme/base/80-number.csc | 388 |
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))))) |
