futhask-base-0.1.0.0: src/Futhask/Array/Element.hs
{-# LANGUAGE TypeFamilyDependencies, FlexibleContexts #-}
{-|
The interface for Haskell arrays, intended to mirror the types of Futhark arrays intuitively.
-}
module Futhask.Array.Element where
import System.IO.Unsafe
import Data.List as L
import Control.Monad (foldM)
-- *Array Type
-- | Defines how an array of something is represented and operated on
class Element elem where
-- |Immutable Array type
type Array elem = arr | arr -> elem
-- |Mutable Array type
type MArray elem = marr | marr -> elem
-- |size of Immutable array
size :: Array elem -> Int
-- |size of mutable array
msize :: MArray elem -> Int
-- |read element without bounds check
unsafeRead :: Array elem -> Int -> elem
-- |peek element without bounds check
unsafePeek :: MArray elem -> Int -> IO elem
-- |poke element without bounds check
unsafePoke :: MArray elem -> Int -> elem -> IO ()
-- |convert array to mutable without copy
unsafeThaw :: Array elem -> IO (MArray elem)
-- |convert array to immutable without copy
unsafeFreeze :: MArray elem -> IO (Array elem)
-- |slice array, new slice points to the same memory
unsafeSlice :: Array elem -> Int -> Int -> Array elem
-- |construct new array from uninitialized memory
scratch :: Int -> IO (MArray elem)
-- | read element, throws error on out of bounds index
read :: Element e => Array e -> Int -> e
read arr idx = if 0 <= idx || idx < size arr then unsafeRead arr idx else error "Index is out of bounds"
-- | read element in the `Maybe` monad
readM :: Element e => Array e -> Int -> Maybe e
readM arr idx = if 0 <= idx || idx < size arr then Just (unsafeRead arr idx) else Nothing
-- | peek element, throws error on out of bounds index
peek :: Element e => MArray e -> Int -> IO e
peek arr idx = if 0 <= idx || idx < msize arr then unsafePeek arr idx else error "Index is out of bounds"
-- | peek element in the `Maybe` monad
peekM :: Element e => MArray e -> Int -> Maybe (IO e)
peekM arr idx = if 0 <= idx || idx < msize arr then Just (unsafePeek arr idx) else Nothing
-- | poke element, throws error on out of bounds index
poke :: Element e => MArray e -> Int -> e -> IO ()
poke arr idx elem = if 0 <= idx || idx < msize arr then unsafePoke arr idx elem else error "Index is out of bounds"
-- | poke element in the `Maybe` monad
pokeM :: Element e => MArray e -> Int -> e -> Maybe (IO ())
pokeM arr idx elem = if 0 <= idx || idx < msize arr then Just (unsafePoke arr idx elem) else Nothing
-- | index operator
(!) :: Element e => Array e -> Int -> e
(!) = Futhask.Array.Element.read
-- | index operator in the `Maybe` monad
(!?) :: Element e => Array e -> Int -> Maybe e
(!?) = Futhask.Array.Element.readM
-- | List of valid indices
indices arr = [0..size arr-1]
-- | List of valid indices for mutable array
mindices arr = [0..msize arr-1]
-- *Mutation
copy :: (Element elem) => MArray elem -> IO (MArray elem)
copy arr = do
arr' <- scratch (msize arr)
mapM_ (\idx -> unsafePeek arr idx >>= unsafePoke arr' idx) (mindices arr)
pure arr'
thaw :: Element elem => Array elem -> IO (MArray elem)
thaw a = unsafeThaw a >>= copy
freeze :: Element elem => MArray elem -> IO (Array elem)
freeze a = copy a >>= unsafeFreeze
-- *Reshaping
unflatten :: (Element elem, Element (Array elem)) => Int -> Int -> Array elem -> Array (Array elem)
unflatten m n arr = if size arr == m*n
then tabulate m (\i -> unsafeSlice arr (n*i) n)
else error "size of array does not fit new dimensions"
unflatten_3d :: (Element elem, Element (Array elem), Element (Array (Array elem))) => Int -> Int -> Int -> Array elem -> Array (Array (Array elem))
unflatten_3d m n l arr = if size arr == m*n*l
then tabulate m (\i -> let o = i*n in tabulate n (\j -> unsafeSlice arr (o+l*j) l))
else error "size of array does not fit new dimensions"
flatten :: (Element elem, Element (Array elem)) => Array (Array elem) -> Array elem
flatten arr = fromList $ concatMap toList (toList arr)
flatten_3d :: (Element elem, Element (Array elem), Element (Array (Array elem))) => Array (Array (Array elem)) -> Array elem
flatten_3d arr = fromList $ concatMap (concatMap toList . toList) (toList arr)
-- *List conversion
toList :: (Element a) => Array a -> [a]
toList = Futhask.Array.Element.foldr (:) []
fromList :: (Element a) => [a] -> Array a
fromList list = unsafePerformIO $ do
arr <- scratch (length list)
mapM_ (\(i, v) -> unsafePoke arr i v) (zip [0..] list)
unsafeFreeze arr
-- *Misc Ops
foldl fun acc arr =
L.foldl (\acc -> fun acc . unsafeRead arr) acc (indices arr)
foldr fun acc arr =
L.foldr (\i acc -> fun (unsafeRead arr i) acc) acc (indices arr)
scatter :: Element a => Array a -> [Int] -> [a] -> Array a
scatter dest is vs = unsafePerformIO $ do
dest <- thaw dest
mapM_ (\(i, v) -> unsafePoke dest i v) (zip is vs)
unsafeFreeze dest
scan :: Element a => (a -> a -> a) -> a -> Array a -> Array a
scan op neu arr = unsafePerformIO $ do
arr' <- scratch (size arr)
foldM (\acc i -> let acc = op acc (unsafeRead arr i) in unsafePoke arr' i acc >> pure acc) neu (indices arr)
unsafeFreeze arr'
tabulate :: Element elem => Int -> (Int -> elem) -> Array elem
tabulate sz f = unsafePerformIO $ do
xs <- scratch sz
mapM_ (\idx -> unsafePoke xs idx (f idx)) [0..sz-1]
unsafeFreeze xs
replicate sz a = tabulate sz (\_ -> a)
map :: (Element a, Element b) => (a -> b) -> Array a -> Array b
map f xs = tabulate (size xs) (f . unsafeRead xs)
map2 :: (Element e0, Element e1, Element e2) => (e0 -> e1 -> e2) -> Array e0 -> Array e1 -> Array e2
map2 f a0 a1 = if size a0 == size a1
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx))
else error "sizes of arguments do not match"
map3 :: (Element e0, Element e1, Element e2, Element e3) => (e0 -> e1 -> e2 -> e3) -> Array e0 -> Array e1 -> Array e2 -> Array e3
map3 f a0 a1 a2 = if size a0 == size a1 && size a1 == size a2
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx))
else error "sizes of arguments do not match"
map4 :: (Element e0, Element e1, Element e2, Element e3, Element e4) => (e0 -> e1 -> e2 -> e3 -> e4) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4
map4 f a0 a1 a2 a3 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx))
else error "sizes of arguments do not match"
map5 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5
map5 f a0 a1 a2 a3 a4 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx))
else error "sizes of arguments do not match"
map6 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6
map6 f a0 a1 a2 a3 a4 a5 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx))
else error "sizes of arguments do not match"
map7 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7
map7 f a0 a1 a2 a3 a4 a5 a6 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx))
else error "sizes of arguments do not match"
map8 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8
map8 f a0 a1 a2 a3 a4 a5 a6 a7 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx))
else error "sizes of arguments do not match"
map9 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9
map9 f a0 a1 a2 a3 a4 a5 a6 a7 a8 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx))
else error "sizes of arguments do not match"
map10 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10
map10 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx))
else error "sizes of arguments do not match"
map11 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10, Element e11) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10 -> e11) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10 -> Array e11
map11 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 a10 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9 && size a9 == size a10
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx) (unsafeRead a10 idx))
else error "sizes of arguments do not match"
map12 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10, Element e11, Element e12) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10 -> e11 -> e12) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10 -> Array e11 -> Array e12
map12 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 a10 a11 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9 && size a9 == size a10 && size a10 == size a11
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx) (unsafeRead a10 idx) (unsafeRead a11 idx))
else error "sizes of arguments do not match"
map13 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10, Element e11, Element e12, Element e13) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10 -> e11 -> e12 -> e13) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10 -> Array e11 -> Array e12 -> Array e13
map13 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 a10 a11 a12 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9 && size a9 == size a10 && size a10 == size a11 && size a11 == size a12
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx) (unsafeRead a10 idx) (unsafeRead a11 idx) (unsafeRead a12 idx))
else error "sizes of arguments do not match"
map14 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10, Element e11, Element e12, Element e13, Element e14) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10 -> e11 -> e12 -> e13 -> e14) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10 -> Array e11 -> Array e12 -> Array e13 -> Array e14
map14 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 a10 a11 a12 a13 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9 && size a9 == size a10 && size a10 == size a11 && size a11 == size a12 && size a12 == size a13
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx) (unsafeRead a10 idx) (unsafeRead a11 idx) (unsafeRead a12 idx) (unsafeRead a13 idx))
else error "sizes of arguments do not match"
map15 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10, Element e11, Element e12, Element e13, Element e14, Element e15) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10 -> e11 -> e12 -> e13 -> e14 -> e15) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10 -> Array e11 -> Array e12 -> Array e13 -> Array e14 -> Array e15
map15 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 a10 a11 a12 a13 a14 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9 && size a9 == size a10 && size a10 == size a11 && size a11 == size a12 && size a12 == size a13 && size a13 == size a14
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx) (unsafeRead a10 idx) (unsafeRead a11 idx) (unsafeRead a12 idx) (unsafeRead a13 idx) (unsafeRead a14 idx))
else error "sizes of arguments do not match"
map16 :: (Element e0, Element e1, Element e2, Element e3, Element e4, Element e5, Element e6, Element e7, Element e8, Element e9, Element e10, Element e11, Element e12, Element e13, Element e14, Element e15, Element e16) => (e0 -> e1 -> e2 -> e3 -> e4 -> e5 -> e6 -> e7 -> e8 -> e9 -> e10 -> e11 -> e12 -> e13 -> e14 -> e15 -> e16) -> Array e0 -> Array e1 -> Array e2 -> Array e3 -> Array e4 -> Array e5 -> Array e6 -> Array e7 -> Array e8 -> Array e9 -> Array e10 -> Array e11 -> Array e12 -> Array e13 -> Array e14 -> Array e15 -> Array e16
map16 f a0 a1 a2 a3 a4 a5 a6 a7 a8 a9 a10 a11 a12 a13 a14 a15 = if size a0 == size a1 && size a1 == size a2 && size a2 == size a3 && size a3 == size a4 && size a4 == size a5 && size a5 == size a6 && size a6 == size a7 && size a7 == size a8 && size a8 == size a9 && size a9 == size a10 && size a10 == size a11 && size a11 == size a12 && size a12 == size a13 && size a13 == size a14 && size a14 == size a15
then tabulate (size a0) (\idx -> f (unsafeRead a0 idx) (unsafeRead a1 idx) (unsafeRead a2 idx) (unsafeRead a3 idx) (unsafeRead a4 idx) (unsafeRead a5 idx) (unsafeRead a6 idx) (unsafeRead a7 idx) (unsafeRead a8 idx) (unsafeRead a9 idx) (unsafeRead a10 idx) (unsafeRead a11 idx) (unsafeRead a12 idx) (unsafeRead a13 idx) (unsafeRead a14 idx) (unsafeRead a15 idx))
else error "sizes of arguments do not match"