taco 0.1.0.0 → 0.2.0.0
raw patch · 9 files changed
+394/−250 lines, 9 filesPVP ok
version bump matches the API change (PVP)
API changes (from Hackage documentation)
- Data.Tensor: Dd :: i -> Dd i
- Data.Tensor: Sd :: Maybe (Vector i) -> Vector i -> i -> Sd i
- Data.Tensor: [D] :: Sh sh -> Dd Int32 -> Sh (sh :# Int32)
- Data.Tensor: [S] :: Sh sh -> Sd Int32 -> Sh (sh :. Int32)
- Data.Tensor: [T] :: Sh i -> Vector a -> Tensor (Sh i) a
- Data.Tensor: [Z] :: Sh Z
- Data.Tensor: [dDim] :: Dd i -> i
- Data.Tensor: [sCml] :: Sd i -> Maybe (Vector i)
- Data.Tensor: [sDim] :: Sd i -> i
- Data.Tensor: [sIdx] :: Sd i -> Vector i
- Data.Tensor: data Sd i
- Data.Tensor: data Sh sh
- Data.Tensor: data Tensor i a
- Data.Tensor: instance GHC.Base.Functor (Data.Tensor.Tensor i)
- Data.Tensor: instance GHC.Classes.Eq a => GHC.Classes.Eq (Data.Tensor.Tensor i a)
- Data.Tensor: instance GHC.Show.Show a => GHC.Show.Show (Data.Tensor.Tensor i a)
- Data.Tensor: newtype Dd i
- Data.Tensor: nnz :: Tensor i a -> Int
- Data.Tensor: shape :: Tensor sh a -> sh
+ Data.Tensor.Compiler: Dd :: i -> Dd i
+ Data.Tensor.Compiler: IncompatIx :: String -> CException
+ Data.Tensor.Compiler: IncompatShape :: String -> CException
+ Data.Tensor.Compiler: Sd :: Maybe (Vector i) -> Vector i -> i -> Sd i
+ Data.Tensor.Compiler: [D] :: Sh sh -> Dd Int32 -> Sh (sh :# Int32)
+ Data.Tensor.Compiler: [S] :: Sh sh -> Sd Int32 -> Sh (sh :. Int32)
+ Data.Tensor.Compiler: [Tensor] :: Sh i -> Vector a -> Tensor (Sh i) a
+ Data.Tensor.Compiler: [Z] :: Sh Z
+ Data.Tensor.Compiler: [dDim] :: Dd i -> i
+ Data.Tensor.Compiler: [sCml] :: Sd i -> Maybe (Vector i)
+ Data.Tensor.Compiler: [sDim] :: Sd i -> i
+ Data.Tensor.Compiler: [sIdx] :: Sd i -> Vector i
+ Data.Tensor.Compiler: contract :: MonadThrow m => [Int] -> Tensor i a -> Tensor i b -> ([Int] -> Tensor i a -> Tensor i b -> Phoas c) -> m (Phoas c)
+ Data.Tensor.Compiler: data CException
+ Data.Tensor.Compiler: data Phoas a
+ Data.Tensor.Compiler: data Sd i
+ Data.Tensor.Compiler: data Sh sh
+ Data.Tensor.Compiler: data Tensor i a
+ Data.Tensor.Compiler: eval :: Phoas a -> a
+ Data.Tensor.Compiler: instance GHC.Classes.Eq Data.Tensor.Compiler.CException
+ Data.Tensor.Compiler: instance GHC.Exception.Exception Data.Tensor.Compiler.CException
+ Data.Tensor.Compiler: instance GHC.Show.Show Data.Tensor.Compiler.CException
+ Data.Tensor.Compiler: let2_ :: Phoas a -> Phoas b -> (a -> b -> Phoas c) -> Phoas c
+ Data.Tensor.Compiler: let_ :: Phoas a -> (a -> Phoas b) -> Phoas b
+ Data.Tensor.Compiler: newtype Dd i
+ Data.Tensor.Compiler: var :: a -> Phoas a
Files
- README.md +3/−1
- src/Data/Dim.hs +3/−0
- src/Data/Shape.hs +13/−6
- src/Data/Shape/Static.hs +11/−0
- src/Data/Tensor.hs +36/−90
- src/Data/Tensor/Compiler.hs +89/−0
- src/Data/Tensor/Compiler/PHOAS.hs +232/−0
- src/Data/TensorTest.hs +0/−148
- taco.cabal +7/−5
README.md view
@@ -1,6 +1,8 @@ # taco -Haskell implementation of the Tensor Algebra COmpiler (`taco`) [1]+Tensor Algebra COmpiler (`taco`).++This library is loosely inspired by [1] but follows an independent design. **NB: Experimental and not yet ready for general use**
src/Data/Dim.hs view
@@ -12,8 +12,11 @@ -- | To define a /sparse/ dimension we need a cumulative array, an index array and a dimensionality parameter data Sd i = Sd {+ -- | Cumulative array (# nonzero entries per degree of freedom). Not all storage formats (e.g. COO for rank-2 tensors) need this information. sCml :: Maybe (V.Vector i)+ -- | Index array (indices of nonzero entries) , sIdx :: V.Vector i+ -- | Size of the tensor along this dimension , sDim :: i } deriving (Eq, Show)
src/Data/Shape.hs view
@@ -1,5 +1,4 @@ {-# language GADTs, TypeOperators #-}- module Data.Shape where import Data.Monoid@@ -8,16 +7,17 @@ -- import GHC.Natural import qualified Data.Vector.Unboxed as VU+import GHC.TypeLits (Nat) import qualified Data.Dim as Dim-+ data Z data sh :# e -- dense data sh :. e -- sparse -- | A statically-typed tensor shape parameter that supports both sparse and dense dimensions.--- Dimensions are indexed with Int32 indices, which should be enough for most applications.+-- Dimensions are indexed with 'Int32' indices, which should be enough for most applications. data Sh sh where Z :: Sh Z -- | Constructor for a dense dimension @@ -26,6 +26,7 @@ S :: Sh sh -> Dim.Sd Int32 -> Sh (sh :. Int32) type D1 = Z :# Int32+type S1 = Z :. Int32 type D2 = (Z :# Int32) :# Int32 type CSR = (Z :# Int32) :. Int32 type COO = (Z :. Int32) :. Int32@@ -49,14 +50,20 @@ rank (S sh _) = 1 + rank sh -- | Dimension of a shape (i.e. list of dimension sizes)-dim :: Sh sh -> [Integer]+dim :: Sh sh -> [Int] dim Z = []-dim (D sh (Dim.Dd m)) = toInteger m : dim sh-dim (S sh (Dim.Sd _ _ m)) = toInteger m : dim sh+dim (D sh (Dim.Dd m)) = fromIntegral m : dim sh+dim (S sh (Dim.Sd _ _ m)) = fromIntegral m : dim sh +-- | Shape of a dense vector+mkD1 :: Int32 -> Sh D1+mkD1 m = Z `D` Dim.Dd m +-- | Shape of a sparse vector+mkS1 :: Int32 -> VU.Vector Int32 -> VU.Vector Int32 -> Sh S1+mkS1 m segv ixv = Z `S` Dim.Sd (Just segv) ixv m -- | Shape of a dense rank-2 tensor (a matrix) mkD2 :: Int32 -> Int32 -> Sh D2
+ src/Data/Shape/Static.hs view
@@ -0,0 +1,11 @@+{-# language GADTs, TypeOperators, DataKinds, KindSignatures #-}+module Data.Shape.Static where++import GHC.TypeLits (Nat)++-- from https://hackage.haskell.org/package/dimensions-0.3.2.0/docs/src/Numeric-Dimensions-Idx.html#Idx+data Idx (ds :: [Nat]) where+ -- | Zero-rank dimensionality - scalar+ Zi :: Idx '[]+ -- | List-like concatenation of indices+ (:!) :: {-# UNPACK #-} !Int -> !(Idx ds) -> Idx (d ': ds)
src/Data/Tensor.hs view
@@ -1,9 +1,12 @@ {-# language GADTs #-} {-# language DeriveFunctor #-} {-# language TypeOperators #-}+{-# language PackageImports #-}+ module Data.Tensor ( -- * Tensor type- Tensor(..), shape, nnz,+ Tensor(..),+ tshape, tdata, nnz, rank, dim, -- * Shape type Sh(..), -- * Dimension types@@ -12,52 +15,29 @@ import qualified Data.Vector as V -- import qualified Data.Vector.Unboxed as VU --- import Data.Word (Word32, Word64) -- import Data.Int (Int32) -import Data.Shape (Sh(..), dim, rank,- Z,- D1, D2, CSR, COO, mkD2, mkCSR, mkCOO)-import qualified Data.Dim as Dim----{- |-IN: Tensor reduction syntax (Einstein notation)--OUT: stride program (how to read/write memory)+import Control.Applicative -taco compiles a tensor expression (e.g. C = A_{ijk}B_{k} ) into a series of nested loops.--dimensions : can be either dense or sparse--internally, tensor data is stored in /dense/ vectors--"contract A_{ijk}B_{k} over the third index"---}+import qualified Data.Shape as Shape (dim, rank)+import Data.Shape (Sh(..), + Z,+ D1, D2, CSR, COO, mkD2, mkCSR, mkCOO) +import qualified Data.Dim as Dim --- | A generic tensor type, polymorphic in the container type as well-data GTensor c i a where- GTensor :: Sh i -> c a -> GTensor c (Sh i) a- -mkGT :: Sh i -> c a -> GTensor c (Sh i) a-mkGT = GTensor-- -- | The 'Tensor' type. Tensor data entries are stored as one single array data Tensor i a where- T :: Sh i -> V.Vector a -> Tensor (Sh i) a + Tensor :: Sh i -> V.Vector a -> Tensor (Sh i) a -- | Construct a tensor given a shape and a vector of entries mkT :: Sh i -> V.Vector a -> Tensor (Sh i) a-mkT = T+mkT = Tensor instance Functor (Tensor i) where- fmap f (T sh v) = T sh (f <$> v)+ fmap f (Tensor sh v) = Tensor sh (f <$> v) -- liftA2' :: (a -> a -> b) -> Tensor i a -> Tensor i a -> Tensor i a -- liftA2' f (T sh1 v1) (T sh2 v2) = mkT sh1 (V.zipWith f v1 v2)@@ -66,81 +46,47 @@ pure' = mkT Z . V.singleton instance (Eq a) => Eq (Tensor i a) where- (T sh1 d1) == (T sh2 d2) = sh1 == sh2 && d1 == d2+ (Tensor sh1 d1) == (Tensor sh2 d2) = sh1 == sh2 && d1 == d2 instance (Show a) => Show (Tensor i a) where- show (T sh d) = unwords [show sh, show $ V.take 5 d, "..."]+ show (Tensor sh d) = unwords [show sh, show $ V.take 5 d, "..."] --- | Access the shape of a tensor-shape :: Tensor sh a -> sh-shape (T sh _) = sh+-- | Access the shape of a 'Tensor'+tshape :: Tensor sh a -> sh+tshape (Tensor sh _) = sh +-- | Access the raw data of a 'Tensor'+tdata :: Tensor sh a -> V.Vector a+tdata (Tensor _ td) = td+ -- | Number of nonzero tensor elements nnz :: Tensor i a -> Int-nnz (T _ td) = V.length td+nnz (Tensor _ td) = V.length td +-- | Tensor rank+rank :: Tensor i a -> Int+rank (Tensor sh _) = Shape.rank sh +-- | Tensor dimensions+dim :: Tensor i a -> [Int]+dim (Tensor sh _) = Shape.dim sh --- * A possible abstract syntax --- data Index i where--- I1 :: i -> Index i--- I2 :: i -> i -> Index (i, i) --- -- | Expressions with tensor operands, e.g. "contract A_{ijk}B_{k} over the third index"---- -- mkConstE = Const <$> mkT---- data Expr a where--- Const :: Tensor (Sh i) a -> Expr (Tensor (Sh i) a)---- data Expr i a where--- -- Const :: a -> Expr a--- Contract :: Index i -> Expr (Sh i) a -> Expr (Sh i) a -> Expr (Sh i) a--- -- (:*:) :: Expr a -> Expr a -> Expr a--- -- (:+:) :: Expr a -> Expr a -> Expr a---- eval (Const x) = x--- eval (Contract ixs a b) = undefined------ data Expr a =--- Const a--- | Contract Int (Expr a) (Expr a)--- -- | Expr a :+: Expr a--- -- | Expr a :*: Expr a--- -- | Expr a :-: Expr a--- -- | Expr a :/: Expr a--- deriving (Eq, Show)---- -- | trivial recursive evaluation function--- eval :: Num t => Expr t -> t--- eval (Const x) = x--- eval (a :+: b) = eval a + eval b--- eval (a :*: b) = eval a * eval b------- | GADT syntax---- data Expr a where--- Const :: a -> Expr a --- -- ^ Sum (elementwise) two expressions--- (:+:) :: Expr a -> Expr a -> Expr a--- -- ^ Multiply (elementwise) two expressions--- (:*:) :: Expr a -> Expr a -> Expr a--- -- ^ Subtract (elementwise) two expressions--- (:-:) :: Expr a -> Expr a -> Expr a-+-- | playground, for future use +-- -- | A generic tensor type, polymorphic in the container type as well+-- data GTensor c i a where+-- GTensor :: Sh i -> c a -> GTensor c (Sh i) a+ +-- mkGT :: Sh i -> c a -> GTensor c (Sh i) a+-- mkGT = GTensor
+ src/Data/Tensor/Compiler.hs view
@@ -0,0 +1,89 @@+{-# language GADTs #-}+{-# language PackageImports #-}++module Data.Tensor.Compiler (+ contract+ -- * Tensor types+ , Tensor(..), Sh(..), Dd(..), Sd(..)+ -- * Syntax+ , Phoas, eval, var, let_, let2_+ -- * Exceptions+ , CException (..)+ )where++import Data.Typeable+import "exceptions" Control.Monad.Catch (MonadThrow(..), throwM, MonadCatch(..), catch)+import Control.Exception (Exception(..))++import Control.Applicative (liftA2, (<|>))++import Data.Tensor (Tensor(..), Sh(..), Dd(..), Sd(..), tshape, tdata, nnz, rank, dim)+import Data.Tensor.Compiler.PHOAS (Phoas(..), let_, let2_, var, lift1, lift2, eval)+++{- |+IN: Tensor reduction syntax (Einstein notation)++OUT: stride program (how to read/write memory)+++taco compiles a tensor expression (e.g. C = A_{ijk}B_{k} ) into a series of nested loops.++dimensions : can be either dense or sparse++internally, tensor data is stored in /dense/ vectors++"contract A_{ijk}B_{k} over the third index"++-}+++++mkVar :: MonadThrow m => [Int] -> Tensor i a -> m (Phoas (Tensor i a))+mkVar ixs0 t = do+ ixs <- mkIxs ixs0 (rank t)+ return $ var t+++mkIxs :: MonadThrow m => [Int] -> Int -> m [Int]+mkIxs ixs mm = go ixs []+ where+ go [] acc = pure acc+ go (i:is) acc | i < 0 =+ throwM $ IncompatIx "Index must be non-negative"+ | i > mm - 1 =+ throwM $ IncompatIx $ unwords ["Index must be smaller than", show mm]+ | otherwise = go is (i : acc)++-- | Tensor contraction+--+-- Inject two 'Tensor' constant into 'Var's, while ensuring that all the contraction indices are compatible with those of the tensors.+--+-- Throws a 'CException' if any index is nonnegative or too large for the shape of the given tensor.+contract :: MonadThrow m =>+ [Int] -- ^ Tensor contraction indices+ -> Tensor i a+ -> Tensor i b+ -> ([Int] -> Tensor i a -> Tensor i b -> Phoas c) -- ^ Contraction function+ -> m (Phoas c)+contract ixs0 t1 t2 f = do+ _ <- mkIxs ixs0 (rank t1)+ ixs <- mkIxs ixs0 (rank t2)+ pure $ let_ (var ixs) $ \ixs' ->+ let2_ (var t1) (var t2) (f ixs')+++-- | Exceptions+data CException = IncompatShape String | IncompatIx String deriving (Eq, Typeable)+instance Show CException where+ show c = case c of+ IncompatShape str -> unwords ["Incompatible shape:", str]+ IncompatIx str -> unwords ["Incompatible index:", str]+instance Exception CException where++++++
+ src/Data/Tensor/Compiler/PHOAS.hs view
@@ -0,0 +1,232 @@+{-# language GADTs, RankNTypes #-}+module Data.Tensor.Compiler.PHOAS where+++-- | Parametric higher-order abstract syntax (PHOAS), after B. Oliveira, A. Loeh, `Abstract Syntax Graphs for Domain Specific Languages` +data Phoas a where+ Var :: a -> Phoas a+ Let :: Phoas a -> (a -> Phoas b) -> Phoas b+ -- Let :: Phoas a -> (a -> Phoas a) -> Phoas a++-- | Inject a constant into the abstract syntax+var :: a -> Phoas a+var = Var++-- | Bind a variable into a closure+let_ :: Phoas a -> (a -> Phoas b) -> Phoas b+let_ = Let++-- | Bind two variables into a closure+let2_ :: Phoas a -> Phoas b -> (a -> b -> Phoas c) -> Phoas c+let2_ a b f = let_ a $ \xa ->+ let_ b $ \xb -> f xa xb++-- letP_ :: Phoas a -> (Phoas a -> Phoas a) -> Phoas a+-- letP_ e f = let_ e (f . Var)++-- letP2_ :: Phoas a -> Phoas a -> (Phoas a -> Phoas a -> Phoas a) -> Phoas a+-- letP2_ a b f = let2_ a b (\x y -> f (Var x) (Var y))+++-- instance Show a => Show (Phoas a) where+-- show e = case e of+-- Var x -> show x+-- -- Let e f -> unwords++++-- | Helper functions +++lift1 :: (a -> b) -> a -> Phoas b+lift1 f = Var . f -- Lift1++lift2 :: (a -> b -> c) -> a -> b -> Phoas c+lift2 f a b = Var (f a b) -- Lift2+++plus :: Num a => a -> a -> Phoas a+plus = lift2 (+)+++-- | Benchmark: `tree 50` should compute the answer instantly.+--+-- This proves that the PHOAS formulation preserves variable sharing+treeE :: Integer -> Phoas Integer+treeE 0 = Var 1+treeE n = let_ (treeE (n - 1)) $ \a -> a `plus` a +++++-- | Semantic function for evaluation+eval :: Phoas a -> a+eval expr = case expr of+ Var x -> x+ Let e f -> eval (f (eval e))+++-- -- | Semantic function for pretty-printing+-- type ClosedExpr = forall a . Phoas a++-- pprint :: ClosedExpr -> String+-- pprint expr = go expr 0+-- where+-- go :: Phoas String -> Int -> String+-- go (Var x) _ = x+-- go (Let e f) c = unwords ["(let", v, "=", go e (c+1), "in", go (f v) (c+1),")"]+-- where+-- v = "v" ++ show c+++++-- * A possible abstract syntax+++{-|+* inner product of two vectors+* matrix-vector action+* matrix-matrix product+-}++data Expr a =+ Konst a+ | Dot (Expr a) (Expr a)+ deriving (Eq, Show)++k :: a -> Expr a+k = Konst+++++-- data UnOp = Sqrt deriving (Eq, Show)+data BinOp = Add | Sub | Mul | Div deriving (Eq, Show)+evalBinOp :: Fractional a => BinOp -> a -> a -> a+evalBinOp op = case op of+ Add -> (+)+ Sub -> (-)+ Mul -> (*)+ Div -> (/)++data EI a =+ -- CW1 UnOp (Expr a)+ CW2 BinOp (Expr a) (Expr a) deriving (Eq, Show)++v1, v2 :: Expr [Int]+v1 = k [1..5]+v2 = k [3..7]++-- evalEI expr = case expr of+-- CW2 op x y ->+++++++-- data Index i where+-- I1 :: i -> Index i+-- I2 :: i -> i -> Index (i, i)++-- -- | Expressions with tensor operands, e.g. "contract A_{ijk}B_{k} over the third index"++-- -- | User-facing grammar:+-- data Expr a where+-- -- | Introduce a constant in the AST+-- Konst :: a -> Expr a+-- -- | Tensor contraction+-- Contr :: i -> (Expr a -> Expr a -> Expr a) -> Expr a+-- -- | Binary componentwise operation+-- CW2 :: (a -> a -> a) -> Expr a -> Expr a -> Expr a++-- k :: a -> Expr a+-- k = Konst++-- (|*|), (|+|) :: Num a => Expr a -> Expr a -> Expr a+-- (|*|) = CW2 (*) +-- (|+|) = CW2 (+)++-- dot = Contr 1 (|+|)++++++-- -- * PHOAS 2++-- data Phoas a where+-- Const :: a -> Phoas a+-- Lift1 :: (a -> b) -> Phoas (a -> b)+-- Let :: Phoas a -> (Phoas a -> Phoas b) -> Phoas b+-- Lambda :: (Phoas a -> Phoas b) -> Phoas (a -> b)+-- App1 :: Phoas (a -> b) -> (Phoas a -> Phoas b)++-- eval expr = case expr of+-- -- Const x -> x+-- Let e f -> f e++-- -- lift1 :: (a -> b) -> a -> Phoas b+-- -- lift1 f = Const . f++-- -- lift2 :: (t2 -> t1 -> t) -> Phoas (t2 -> t1 -> t)+-- -- lift2 f = Const $ \a b -> f a b +++-- -- +++-- data Phoas a =+-- Lit Int+-- -- | Lift1 (a -> a) (a -> Phoas a)+-- -- | Add (Phoas a) (Phoas a)+-- | Let (Phoas a) (a -> Phoas a)+-- | Let2 (Phoas a) (Phoas a) (a -> a -> Phoas a)+-- | Var a++-- evalPhoas expr = case expr of+-- Var x -> x+-- Lit i -> i+-- -- Add e0 e1 -> evalPhoas e0 + evalPhoas e1+-- Let e f -> evalPhoas $ f e' where e' = evalPhoas e+-- Let2 e0 e1 f -> evalPhoas $ f e0' e1' where+-- e0' = evalPhoas e0+-- e1' = evalPhoas e1+++-- contract ixs (T sh)++-- eval (Const x) = x+-- eval (Contract ixs a b) = undefined++++-- data Expr a =+-- Const a+-- | Contract Int (Expr a) (Expr a)+-- -- | Expr a :+: Expr a+-- -- | Expr a :*: Expr a+-- -- | Expr a :-: Expr a+-- -- | Expr a :/: Expr a+-- deriving (Eq, Show)++-- -- | trivial recursive evaluation function+-- eval :: Num t => Expr t -> t+-- eval (Const x) = x+-- eval (a :+: b) = eval a + eval b+-- eval (a :*: b) = eval a * eval b+++++-- | GADT syntax++-- data Expr a where+-- Const :: a -> Expr a +-- -- ^ Sum (elementwise) two expressions+-- (:+:) :: Expr a -> Expr a -> Expr a+-- -- ^ Multiply (elementwise) two expressions+-- (:*:) :: Expr a -> Expr a -> Expr a+-- -- ^ Subtract (elementwise) two expressions+-- (:-:) :: Expr a -> Expr a -> Expr a
− src/Data/TensorTest.hs
@@ -1,148 +0,0 @@-{-# language MultiParamTypeClasses, TypeFamilies, FlexibleContexts #-}-{-# language TypeOperators #-}-{-# language GADTs #-}-module Data.TensorTest where--import Data.List (splitAt, unfoldr)--import qualified Data.Vector.Unboxed as V--- import Data.Tensor---data Exp a where- Lift :: (a -> b) -> Exp a -> Exp b----- | tensor elements may be indexed-class Ord (Ix a) => Elem a where- type Ix a :: *- type Ev a :: *- proj :: a -> (Ix a, Ev a)- prod :: Ix a -> Ev a -> a---data E1 a = E1 Int a deriving Show--instance Elem (E1 a) where- type Ix (E1 a) = Int- type Ev (E1 a) = a- proj (E1 i x) = (i, x)- prod i x = E1 i x--data E2 a = E2 Int Int a deriving Show--instance Elem (E2 a) where- type Ix (E2 a) = (Int, Int)- type Ev (E2 a) = a- proj (E2 i j x) = ((i, j), x)- prod (i, j) x = E2 i j x----- -----spUnion' :: Elem a => (Ev a -> Ev a -> Ev a) -> [a] -> [a] -> [a]-spUnion' ff = go where- go [] y = y- go x [] = x- go xv@(elx:xs) yv@(ely:ys) =- let- (ix, elvx) = proj elx- (iy, elvy) = proj ely- in - case compare ix iy of- EQ -> (prod ix (ff elvx elvy)) : go xs ys- LT -> prod ix elvx : go xs yv- GT -> prod iy elvy : go xv ys- ------- | sparse vector /union/ as 2-way merge-spUnion :: Ord i => (a -> a -> a) -> [(i, a)] -> [(i, a)] -> [(i, a)]-spUnion ff = go where- go [] y = y- go x [] = x- go xv@((i,x):xs) yv@((j,y):ys) = - case compare i j of EQ -> (i, ff x y) : go xs ys- LT -> (i, x) : go xs yv- GT -> (j, y) : go xv ys---- | sparse vector /component-wise intersection/ as 2-way merge-spIntersect :: Ord i => (a -> a -> a) -> [(i, a)] -> [(i, a)] -> [(i, a)]-spIntersect gg = go where- go [] _ = []- go _ [] = []- go xv@((i,x):xs) yv@((j,y):ys) =- case compare i j of EQ -> (i, gg x y) : go xs ys- LT -> go xs yv- GT -> go xv ys--spAdd :: (Num a, Ord i) => [(i, a)] -> [(i, a)] -> [(i, a)] -spAdd = spUnion (+)--spMul :: (Num a, Ord i) => [(i, a)] -> [(i, a)] -> [(i, a)]-spMul = spIntersect (*)---v0, v1 :: [(Int, Int)]-v0 = [(0, 1), (2, 2), (5, 1)]-v1 = [(0, 2), (1, 3), (2, 3), (4, 1)]------ | "AND"-conjunction :: (a -> b -> c) -> Maybe a -> Maybe b -> Maybe c-conjunction f a b = case (a, b) of- (Just x, Just y) -> Just $ f x y- _ -> Nothing---- | "OR"-disjunction :: (t -> t -> t) -> Maybe t -> Maybe t -> Maybe t-disjunction f a b = case (a, b) of- (Nothing, Nothing) -> Nothing- (Just x, Nothing) -> Just x- (Nothing, Just y) -> Just y- (Just x, Just y) -> Just $ f x y--sumMaybe :: Num a => Maybe a -> Maybe a -> Maybe a -sumMaybe = disjunction (+)--prodMaybe :: Num a => Maybe a -> Maybe a -> Maybe a -prodMaybe = conjunction (*)-------- chunk sparse lists according to element index---- chunkBy q ll = go ll []--- where--- go _ acc = acc--- go (e:es) acc--- | q e = go es (e : acc)--- -- | otherwise =--chunks :: Int -> [a] -> [[a]]-chunks _ [] = []-chunks n ll = h : chunks n t where- (h, t) = splitAt n ll--chunksWhile :: (a -> Bool) -> [a] -> [[a]]-chunksWhile _ [] = []-chunksWhile q ll = h : chunksWhile q t where- (h, t) = (takeWhile q ll, dropWhile q ll)--chunksWhile' :: (a -> Bool) -> [a] -> [[a]]-chunksWhile' q = unfoldr genf where- genf ll =- if null h- then Nothing- else Just (h, drop (length h) ll)- where h = takeWhile q ll
taco.cabal view
@@ -1,6 +1,6 @@ name: taco-version: 0.1.0.0-synopsis: Haskell port of the Tensor Algebra COmpiler+version: 0.2.0.0+synopsis: Tensor Algebra COmpiler description: This library provides types and a compiler for tensor expressions. homepage: https://github.com/ocramz/taco-hs#readme license: BSD3@@ -16,10 +16,12 @@ library hs-source-dirs: src ghc-options: -Wall- exposed-modules: Data.Tensor- other-modules: Data.Shape+ exposed-modules: Data.Tensor.Compiler+ other-modules: Data.Tensor+ Data.Shape Data.Dim- Data.TensorTest+ Data.Shape.Static+ Data.Tensor.Compiler.PHOAS build-depends: base >= 4.7 && < 5 , exceptions , mtl