packages feed

atrophy-0.2.0.0: src/Atrophy/Internal/Prim.hs

{-# LANGUAGE CPP #-}
{-# LANGUAGE MagicHash #-}
{-# LANGUAGE UnboxedTuples #-}
{-# LANGUAGE AllowAmbiguousTypes #-}
{-# OPTIONS_HADDOCK hide #-}

#include "MachDeps.h"

-- | Double-word and multi-limb primitives. Everything in here is @INLINE@ and
-- is expected to compile down to a handful of instructions. Limbs are passed
-- most significant first, results are returned as unboxed tuples so that
-- nothing is ever allocated.
module Atrophy.Internal.Prim
  ( -- * 64-bit
    mulHi64
  , mulFull64
  , quotRem128By64
  , quotRemWord64
  , quotWord32
  , ltW
  , addCarry64
  , subBorrow64
  , adc64
    -- * 128-bit, as pairs of limbs
  , add128
  , sub128
  , shr128
  , shr128Small
  , shl128Small
  , clz128
  , mulLo128
  , mulHi128
  , mulHi128By64
  , div3By2
    -- * Type level
  , natWord64
  , natInt
  ) where

import Data.Bits
import GHC.Exts
import GHC.TypeNats (KnownNat, natVal')
import GHC.Word

#if WORD_SIZE_IN_BITS == 64

-- | High 64 bits of the 128-bit product. A single @mul@.
{-# INLINE mulHi64 #-}
mulHi64 :: Word64 -> Word64 -> Word64
mulHi64 (W64# a) (W64# b) = case timesWord2# (word64ToWord# a) (word64ToWord# b) of
  (# h, _ #) -> W64# (wordToWord64# h)

-- | @(# hi, lo #)@ of the 128-bit product. A single @mul@.
{-# INLINE mulFull64 #-}
mulFull64 :: Word64 -> Word64 -> (# Word64, Word64 #)
mulFull64 (W64# a) (W64# b) = case timesWord2# (word64ToWord# a) (word64ToWord# b) of
  (# h, l #) -> (# W64# (wordToWord64# h), W64# (wordToWord64# l) #)

-- | @quotRem128By64 hi lo d@ divides @hi * 2^64 + lo@ by @d@. Requires @hi < d@.
-- A single @div@.
{-# INLINE quotRem128By64 #-}
quotRem128By64 :: Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
quotRem128By64 (W64# h) (W64# l) (W64# d) =
  case quotRemWord2# (word64ToWord# h) (word64ToWord# l) (word64ToWord# d) of
    (# q, r #) -> (# W64# (wordToWord64# q), W64# (wordToWord64# r) #)

-- | Unchecked division, the divisor must not be zero.
{-# INLINE quotRemWord64 #-}
quotRemWord64 :: Word64 -> Word64 -> (# Word64, Word64 #)
quotRemWord64 (W64# n) (W64# d) = case quotRemWord# (word64ToWord# n) (word64ToWord# d) of
  (# q, r #) -> (# W64# (wordToWord64# q), W64# (wordToWord64# r) #)

-- | 1 if @a < b@, 0 otherwise. Branchless.
{-# INLINE ltW #-}
ltW :: Word64 -> Word64 -> Word64
ltW (W64# a) (W64# b) = W64# (wordToWord64# (int2Word# (ltWord# (word64ToWord# a) (word64ToWord# b))))

#else

{-# INLINE mulHi64 #-}
mulHi64 :: Word64 -> Word64 -> Word64
mulHi64 a b = case mulFull64 a b of (# h, _ #) -> h

{-# INLINE mulFull64 #-}
mulFull64 :: Word64 -> Word64 -> (# Word64, Word64 #)
mulFull64 a b =
  let !aL = a .&. 0xffffffff
      !aH = a `unsafeShiftR` 32
      !bL = b .&. 0xffffffff
      !bH = b `unsafeShiftR` 32
      !ll = aL * bL
      !lh = aL * bH
      !hl = aH * bL
      !hh = aH * bH
      !mid = (ll `unsafeShiftR` 32) + (lh .&. 0xffffffff) + (hl .&. 0xffffffff)
  in (# hh + (lh `unsafeShiftR` 32) + (hl `unsafeShiftR` 32) + (mid `unsafeShiftR` 32), a * b #)

{-# INLINE quotRem128By64 #-}
quotRem128By64 :: Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
quotRem128By64 h l d =
  case (toInteger h `unsafeShiftL` 64 .|. toInteger l) `quotRem` toInteger d of
    (q, r) -> (# fromInteger q, fromInteger r #)

{-# INLINE quotRemWord64 #-}
quotRemWord64 :: Word64 -> Word64 -> (# Word64, Word64 #)
quotRemWord64 n d = case quotRem n d of (q, r) -> (# q, r #)

{-# INLINE ltW #-}
ltW :: Word64 -> Word64 -> Word64
ltW a b = if a < b then 1 else 0

#endif

-- | Unchecked division, the divisor must not be zero.
{-# INLINE quotWord32 #-}
quotWord32 :: Word32 -> Word32 -> Word32
quotWord32 (W32# n) (W32# d) = W32# (quotWord32# n d)

-- | @(# sum, carry #)@
{-# INLINE addCarry64 #-}
addCarry64 :: Word64 -> Word64 -> (# Word64, Word64 #)
addCarry64 a b = let !s = a + b in (# s, ltW s a #)

-- | @(# difference, borrow #)@
{-# INLINE subBorrow64 #-}
subBorrow64 :: Word64 -> Word64 -> (# Word64, Word64 #)
subBorrow64 a b = (# a - b, ltW a b #)

-- | @a + b + carry@, returning @(# sum, carry #)@. @carry@ must be 0 or 1.
{-# INLINE adc64 #-}
adc64 :: Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
adc64 a b c =
  let !s1 = a + b
      !s = s1 + c
  in (# s, ltW s1 a + ltW s s1 #)

{-# INLINE add128 #-}
add128 :: Word64 -> Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
add128 a1 a0 b1 b0 = case addCarry64 a0 b0 of (# s, c #) -> (# a1 + b1 + c, s #)

{-# INLINE sub128 #-}
sub128 :: Word64 -> Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
sub128 a1 a0 b1 b0 = case subBorrow64 a0 b0 of (# s, c #) -> (# a1 - b1 - c, s #)

-- | Shift right by @0 <= s < 64@, branchless.
{-# INLINE shr128Small #-}
shr128Small :: Word64 -> Word64 -> Int -> (# Word64, Word64 #)
shr128Small x1 x0 s =
  (# x1 `unsafeShiftR` s
  , (x0 `unsafeShiftR` s) .|. ((x1 `unsafeShiftL` 1) `unsafeShiftL` (63 - s))
  #)

-- | Shift left by @0 <= s < 64@, branchless.
{-# INLINE shl128Small #-}
shl128Small :: Word64 -> Word64 -> Int -> (# Word64, Word64 #)
shl128Small x1 x0 s =
  (# (x1 `unsafeShiftL` s) .|. ((x0 `unsafeShiftR` 1) `unsafeShiftR` (63 - s))
  , x0 `unsafeShiftL` s
  #)

-- | Shift right by @0 <= s < 128@.
{-# INLINE shr128 #-}
shr128 :: Word64 -> Word64 -> Int -> (# Word64, Word64 #)
shr128 x1 x0 s
  | s < 64 = shr128Small x1 x0 s
  | otherwise = (# 0, x1 `unsafeShiftR` (s - 64) #)

{-# INLINE clz128 #-}
clz128 :: Word64 -> Word64 -> Int
clz128 x1 x0 = if x1 == 0 then 64 + countLeadingZeros x0 else countLeadingZeros x1

-- | Low 128 bits of the product.
{-# INLINE mulLo128 #-}
mulLo128 :: Word64 -> Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
mulLo128 a1 a0 b1 b0 = case mulFull64 a0 b0 of
  (# h, l #) -> (# h + a0 * b1 + a1 * b0, l #)

-- | High 128 bits of the 256-bit product.
{-# INLINE mulHi128 #-}
mulHi128 :: Word64 -> Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
mulHi128 a1 a0 b1 b0 =
  case mulFull64 a0 b0 of { (# h00, _ #) ->
  case mulFull64 a0 b1 of { (# h01, l01 #) ->
  case mulFull64 a1 b0 of { (# h10, l10 #) ->
  case mulFull64 a1 b1 of { (# h11, l11 #) ->
  case addCarry64 h00 l01 of { (# s1, c1a #) ->
  case addCarry64 s1 l10 of { (# _, c1b #) ->
  case addCarry64 h01 h10 of { (# s2a, c2a #) ->
  case addCarry64 s2a l11 of { (# s2b, c2b #) ->
  case addCarry64 s2b (c1a + c1b) of { (# s2, c2c #) ->
  (# h11 + c2a + c2b + c2c, s2 #) }}}}}}}}}

-- | High 128 bits of the 192-bit product of a 128-bit and a 64-bit number.
{-# INLINE mulHi128By64 #-}
mulHi128By64 :: Word64 -> Word64 -> Word64 -> (# Word64, Word64 #)
mulHi128By64 a1 a0 b =
  case mulFull64 a1 b of { (# h1, l1 #) ->
  case mulHi64 a0 b of { h0 ->
  case addCarry64 l1 h0 of { (# s, c #) ->
  (# h1 + c, s #) }}}

-- | Divide the 192-bit @(u2, u1, u0)@ by the normalized (top bit set) 128-bit
-- @(d1, d0)@, given @(u2, u1) < (d1, d0)@. The quotient fits in 64 bits.
-- Returns @(# quotient, remainder hi, remainder lo #)@.
--
-- Knuth's algorithm D, one step: estimate from the top limbs with a hardware
-- division, then correct at most twice.
{-# INLINE div3By2 #-}
div3By2 :: Word64 -> Word64 -> Word64 -> Word64 -> Word64 -> (# Word64, Word64, Word64 #)
div3By2 u2 u1 u0 d1 d0 =
  let !qhat = if u2 >= d1 then maxBound else case quotRem128By64 u2 u1 d1 of (# q, _ #) -> q
  in
  case mulFull64 qhat d0 of { (# a1, a0 #) ->
  case mulFull64 qhat d1 of { (# b1, b0 #) ->
  case addCarry64 b0 a1 of { (# p1, c #) ->
  correct qhat (b1 + c) p1 a0 }}}
  where
  correct :: Word64 -> Word64 -> Word64 -> Word64 -> (# Word64, Word64, Word64 #)
  correct !q !p2 !p1 !p0
    | p2 > u2 || (p2 == u2 && (p1 > u1 || (p1 == u1 && p0 > u0))) =
        case subBorrow64 p0 d0 of { (# p0', b0 #) ->
        case subBorrow64 p1 d1 of { (# t1, b1a #) ->
        case subBorrow64 t1 b0 of { (# p1', b1b #) ->
        correct (q - 1) (p2 - b1a - b1b) p1' p0' }}}
    | otherwise =
        -- the remainder is below d, so its top limb is zero
        case subBorrow64 u0 p0 of { (# r0, b0 #) ->
        (# q, u1 - p1 - b0, r0 #) }

-- | A type level natural as a 'Word64' literal. Folds at compile time.
{-# INLINE natWord64 #-}
natWord64 :: forall n. KnownNat n => Word64
natWord64 = fromIntegral (natVal' (proxy# :: Proxy# n))

-- | A type level natural as an 'Int' literal. Folds at compile time.
{-# INLINE natInt #-}
natInt :: forall n. KnownNat n => Int
natInt = fromIntegral (natVal' (proxy# :: Proxy# n))