matrix 0.3.1.1 → 0.3.2.0
raw patch · 4 files changed
+241/−126 lines, 4 filesPVP ok
version bump matches the API change (PVP)
API changes (from Hackage documentation)
+ Data.Matrix: multStd2 :: Num a => Matrix a -> Matrix a -> Matrix a
Files
- Data/Matrix.hs +217/−119
- bench/mult.hs +6/−1
- matrix.cabal +1/−1
- test/Main.hs +17/−5
Data/Matrix.hs view
@@ -43,6 +43,7 @@ -- ** Functions , multStd+ , multStd2 , multStrassen , multStrassenMixed -- * Linear transformations@@ -96,18 +97,28 @@ -- @i,j :: Int@, then @m ! (i,j)@ is the element in the @i@-th row and -- @j@-th column of @m@. data Matrix a = M {- nrows :: {-# UNPACK #-} !Int -- ^ Number of rows.- , ncols :: {-# UNPACK #-} !Int -- ^ Number of columns.- , mvect :: V.Vector a -- ^ Content of the matrix as a plain vector.- } deriving Eq+ nrows :: {-# UNPACK #-} !Int -- ^ Number of rows.+ , ncols :: {-# UNPACK #-} !Int -- ^ Number of columns.+ , rowOffset :: {-# UNPACK #-} !Int+ , colOffset :: {-# UNPACK #-} !Int+ , vcols :: {-# UNPACK #-} !Int -- ^ Number of columns of the matrix without offset+ , mvect :: V.Vector a -- ^ Content of the matrix as a plain vector.+ } +instance Eq a => Eq (Matrix a) where+ m1 == m2 =+ let r = nrows m1+ c = ncols m1+ in and $ (r == nrows m2) : (c == ncols m2)+ : [ m1 ! (i,j) == m2 ! (i,j) | i <- [1 .. r] , j <- [1 .. c] ]+ -- | Just a cool way to output the size of a matrix. sizeStr :: Int -> Int -> String sizeStr n m = show n ++ "x" ++ show m -- | Display a matrix as a 'String' using the 'Show' instance of its elements. prettyMatrix :: Show a => Matrix a -> String-prettyMatrix m@(M _ _ v) = unlines+prettyMatrix m@(M _ _ _ _ _ v) = unlines [ "( " <> unwords (fmap (\j -> fill mx $ show $ m ! (i,j)) [1..ncols m]) <> " )" | i <- [1..nrows m] ] where mx = V.maximum $ fmap (length . show) v@@ -117,13 +128,15 @@ show = prettyMatrix instance NFData a => NFData (Matrix a) where- rnf (M _ _ v) = rnf v+ rnf = rnf . mvect --- | /O(rows*cols)/. Similar to 'V.force', drop any extra memory.+-- | /O(rows*cols)/. Similar to 'V.force'. It copies the matrix content+-- dropping any extra memory. ----- Useful when using 'getRow' from a big matrix.+-- Useful when using 'submatrix' from a big matrix.+-- forceMatrix :: Matrix a -> Matrix a-forceMatrix (M n m v) = M n m $ V.force v+forceMatrix m = matrix (nrows m) (ncols m) $ \(i,j) -> m ! (i,j) ------------------------------------------------------- -------------------------------------------------------@@ -131,7 +144,7 @@ instance Functor Matrix where {-# INLINE fmap #-}- fmap f (M n m v) = M n m $ V.map f v+ fmap f (M n m ro co w v) = M n m ro co w $ V.map f v ------------------------------------------------------- -------------------------------------------------------@@ -146,9 +159,11 @@ mapRow :: (Int -> a -> a) -- ^ Function takes the current column as additional argument. -> Int -- ^ Row to map. -> Matrix a -> Matrix a-mapRow f r (M n m v) =- M n m $ V.imap (\k x -> let (i,j) = decode m k- in if i == r then f j x else x) v+mapRow f r0 (M n m ro co w v) =+ let r = r0 + ro+ in M n m ro co w $+ V.imap (\k x -> let (i,j) = decode w k+ in if i == r then f j x else x) v -- | /O(rows*cols)/. Map a function over a column. -- Example:@@ -160,19 +175,21 @@ mapCol :: (Int -> a -> a) -- ^ Function takes the current row as additional argument. -> Int -- ^ Column to map. -> Matrix a -> Matrix a-mapCol f c (M n m v) =- M n m $ V.imap (\k x -> let (i,j) = decode m k- in if j == c then f i x else x) v+mapCol f c0 (M n m ro co w v) =+ let c = c0 + co+ in M n m ro co w $+ V.imap (\k x -> let (i,j) = decode w k+ in if j == c then f i x else x) v ------------------------------------------------------- ------------------------------------------------------- ---- FOLDABLE AND TRAVERSABLE INSTANCES instance Foldable Matrix where- foldMap f = foldMap f . mvect+ foldMap f = foldMap f . mvect . forceMatrix instance Traversable Matrix where- sequenceA (M n m v) = fmap (M n m) $ sequenceA v+ sequenceA m = fmap (M (nrows m) (ncols m) 0 0 (ncols m)) . sequenceA . mvect $ forceMatrix m ------------------------------------------------------- -------------------------------------------------------@@ -191,7 +208,8 @@ Int -- ^ Rows -> Int -- ^ Columns -> Matrix a-zero n m = M n m $ V.replicate (n*m) 0+{-# INLINE zero #-}+zero n m = M n m 0 0 m $ V.replicate (n*m) 0 -- | /O(rows*cols)/. Generate a matrix from a generator function. -- Example of usage:@@ -205,7 +223,7 @@ -> ((Int,Int) -> a) -- ^ Generator function -> Matrix a {-# INLINE matrix #-}-matrix n m f = M n m $ V.generate (n*m) $ f . decode m+matrix n m f = M n m 0 0 m $ V.generate (n*m) $ f . decode m -- | /O(rows*cols)/. Identity matrix of the given order. --@@ -233,9 +251,9 @@ -> [a] -- ^ List of elements -> Matrix a {-# INLINE fromList #-}-fromList n m = M n m . V.fromListN (n*m)+fromList n m = M n m 0 0 m . V.fromListN (n*m) --- | Create a matrix from an non-empty list of non-empty lists.+-- | Create a matrix from a non-empty list of non-empty lists. -- /Each list must have at least as many elements as the first list/. -- Examples: --@@ -257,11 +275,13 @@ -- | /O(1)/. Represent a vector as a one row matrix. rowVector :: V.Vector a -> Matrix a-rowVector v = M 1 (V.length v) v+rowVector v = M 1 m 0 0 m v+ where+ m = V.length v -- | /O(1)/. Represent a vector as a one column matrix. colVector :: V.Vector a -> Matrix a-colVector v = M (V.length v) 1 v+colVector v = M (V.length v) 1 0 0 1 v -- | /O(rows*cols)/. Permutation matrix. --@@ -298,15 +318,21 @@ ---- ACCESSING -- | /O(1)/. Get an element of a matrix. Indices range from /(1,1)/ to /(n,m)/.+-- It returns an 'error' if the requested element is outside of range. getElem :: Int -- ^ Row -> Int -- ^ Column -> Matrix a -- ^ Matrix -> a-getElem i j a@(M n m _)- | i > n || j > m || i < 1 || j < 1 =- error $ "Trying to get the " ++ show (i,j) ++ " element from a "- ++ sizeStr n m ++ " matrix."- | otherwise = unsafeGet i j a+{-# INLINE getElem #-}+getElem i j m =+ case safeGet i j m of+ Just x -> x+ Nothing -> error+ $ "getElem: Trying to get the "+ ++ show (i,j)+ ++ " element from a "+ ++ sizeStr (nrows m) (ncols m)+ ++ " matrix." -- | /O(1)/. Unsafe variant of 'getElem', without bounds checking. unsafeGet :: Int -- ^ Row@@ -314,7 +340,7 @@ -> Matrix a -- ^ Matrix -> a {-# INLINE unsafeGet #-}-unsafeGet i j (M _ m v) = V.unsafeIndex v $ encode m (i,j)+unsafeGet i j (M _ _ ro co w v) = V.unsafeIndex v $ encode w (i+ro,j+co) -- | Short alias for 'getElem'. (!) :: Matrix a -> (Int,Int) -> a@@ -323,17 +349,19 @@ -- | Variant of 'getElem' that returns Maybe instead of an error. safeGet :: Int -> Int -> Matrix a -> Maybe a-safeGet i j a@(M n m _)+safeGet i j a@(M n m _ _ _ _) | i > n || j > m || i < 1 || j < 1 = Nothing | otherwise = Just $ unsafeGet i j a -- | /O(1)/. Get a row of a matrix as a vector. getRow :: Int -> Matrix a -> V.Vector a-getRow i (M _ m v) = V.slice (m*(i-1)) m v+{-# INLINE getRow #-}+getRow i (M _ m ro co w v) = V.slice (w*(i-1+ro) + co) m v -- | /O(rows)/. Get a column of a matrix as a vector. getCol :: Int -> Matrix a -> V.Vector a-getCol j (M n m v) = V.generate n $ \i -> v V.! encode m (i+1,j)+{-# INLINE getCol #-}+getCol j (M n _ ro co w v) = V.generate n $ \i -> v V.! encode w (i+1+ro,j+co) -- | /O(min rows cols)/. Diagonal of a /not necessarily square/ matrix. getDiag :: Matrix a -> V.Vector a@@ -341,37 +369,53 @@ where k = min (nrows m) (ncols m) --- | /O(1)/. Transform a 'Matrix' to a 'V.Vector' of size /rows*cols/.+-- | /O(rows*cols)/. Transform a 'Matrix' to a 'V.Vector' of size /rows*cols/. -- This is equivalent to get all the rows of the matrix using 'getRow' -- and then append them, but far more efficient. getMatrixAsVector :: Matrix a -> V.Vector a-getMatrixAsVector = mvect+getMatrixAsVector = mvect . forceMatrix ------------------------------------------------------- ------------------------------------------------------- ---- MANIPULATING MATRICES -msetElem:: PrimMonad m => a -> Int -> (Int,Int) -> MV.MVector (PrimState m) a -> m ()+msetElem :: PrimMonad m+ => a -- ^ New element+ -> Int -- ^ Number of columns of the matrix+ -> Int -- ^ Row offset+ -> Int -- ^ Column offset+ -> (Int,Int) -- ^ Position to set the new element+ -> MV.MVector (PrimState m) a -- ^ Mutable vector+ -> m () {-# INLINE msetElem #-}-msetElem x m p v = MV.write v (encode m p) x+msetElem x w ro co (i,j) v = MV.write v (encode w (i+ro,j+co)) x -unsafeMset:: PrimMonad m => a -> Int -> (Int,Int) -> MV.MVector (PrimState m) a -> m ()+unsafeMset :: PrimMonad m+ => a -- ^ New element+ -> Int -- ^ Number of columns of the matrix+ -> Int -- ^ Row offset+ -> Int -- ^ Column offset+ -> (Int,Int) -- ^ Position to set the new element+ -> MV.MVector (PrimState m) a -- ^ Mutable vector+ -> m () {-# INLINE unsafeMset #-}-unsafeMset x m p v = MV.unsafeWrite v (encode m p) x+unsafeMset x w ro co (i,j) v = MV.unsafeWrite v (encode w (i+ro,j+co)) x -- | Replace the value of a cell in a matrix. setElem :: a -- ^ New value. -> (Int,Int) -- ^ Position to replace. -> Matrix a -- ^ Original matrix. -> Matrix a -- ^ Matrix with the given position replaced with the given value.-setElem x p (M n m v) = M n m $ V.modify (msetElem x m p) v+{-# INLINE setElem #-}+setElem x p (M n m ro co w v) = M n m ro co w $ V.modify (msetElem x w ro co p) v -- | Unsafe variant of 'setElem', without bounds checking. unsafeSet :: a -- ^ New value. -> (Int,Int) -- ^ Position to replace. -> Matrix a -- ^ Original matrix. -> Matrix a -- ^ Matrix with the given position replaced with the given value.-unsafeSet x p (M n m v) = M n m $ V.modify (unsafeMset x m p) v+{-# INLINE unsafeSet #-}+unsafeSet x p (M n m ro co w v) = M n m ro co w $ V.modify (unsafeMset x w ro co p) v -- | /O(rows*cols)/. The transpose of a matrix. -- Example:@@ -409,36 +453,31 @@ -> Int -- ^ Number of columns. -> Matrix a -> Matrix a-setSize e n m a@(M r c _) = M n m $ V.generate (n*m) $- \k -> let (i,j) = d k- in if i <= r && j <= c- then a ! (i,j)- else e- where- d = decode m+{-# INLINE setSize #-}+setSize e n m a@(M n0 m0 _ _ _ _) = matrix n m $ \(i,j) ->+ if i <= n0 && j <= m0+ then unsafeGet i j a+ else e ------------------------------------------------------- ------------------------------------------------------- ---- WORKING WITH BLOCKS --- | /O(subrows*subcols)/. Extract a submatrix given row and column limits.+-- | /O(1)/. Extract a submatrix given row and column limits. -- Example: -- -- > ( 1 2 3 ) -- > ( 4 5 6 ) ( 2 3 ) -- > submatrix 1 2 2 3 ( 7 8 9 ) = ( 5 6 ) submatrix :: Int -- ^ Starting row- -> Int -- ^ Ending row+ -> Int -- ^ Ending row -> Int -- ^ Starting column- -> Int -- ^ Ending column+ -> Int -- ^ Ending column -> Matrix a -> Matrix a {-# INLINE submatrix #-}-submatrix r1 r2 c1 c2 (M _ m vs) = M r' c' $ V.generate (r'*c') $- \k -> let (i,j) = decode c' k in vs V.! encode m (i+r1-1,j+c1-1)- where- r' = r2-r1+1- c' = c2-c1+1+submatrix r1 r2 c1 c2 (M _ _ ro co w v) =+ M (r2-r1+1) (c2-c1+1) (ro+r1-1) (co+c1-1) w v -- | /O(rows*cols)/. Remove a row and a column from a matrix. -- Example:@@ -450,10 +489,12 @@ -> Int -- ^ Column @c@ to remove. -> Matrix a -- ^ Original matrix. -> Matrix a -- ^ Matrix with row @r@ and column @c@ removed.-minorMatrix r c (M n m v) =- M (n-1) (m-1) $ V.ifilter (\k _ -> let (i,j) = decode m k in i /= r && j /= c) v+minorMatrix r0 c0 (M n m ro co w v) =+ let r = r0 + ro+ c = c0 + co+ in M (n-1) (m-1) ro co (w-1) $ V.ifilter (\k _ -> let (i,j) = decode w k in i /= r && j /= c) v --- | /O(rows*cols)/. Make a block-partition of a matrix using a given element as reference.+-- | /O(1)/. Make a block-partition of a matrix using a given element as reference. -- The element will stay in the bottom-right corner of the top-left corner matrix. -- -- > ( ) ( | )@@ -472,14 +513,15 @@ -- -- Where T = Top, B = Bottom, L = Left, R = Right. ----- Implementation is done via slicing of vectors. splitBlocks :: Int -- ^ Row of the splitting element. -> Int -- ^ Column of the splitting element. -> Matrix a -- ^ Matrix to split. -> (Matrix a,Matrix a ,Matrix a,Matrix a) -- ^ (TL,TR,BL,BR)-splitBlocks i j a@(M n m _) = ( submatrix 1 i 1 j a , submatrix 1 i (j+1) m a- , submatrix (i+1) n 1 j a , submatrix (i+1) n (j+1) m a )+{-# INLINE[1] splitBlocks #-}+splitBlocks i j a@(M n m _ _ _ _) =+ ( submatrix 1 i 1 j a , submatrix 1 i (j+1) m a+ , submatrix (i+1) n 1 j a , submatrix (i+1) n (j+1) m a ) -- | Join blocks of the form detailed in 'splitBlocks'. Precisely: --@@ -488,6 +530,7 @@ -- > <-> -- > (bl <|> br) joinBlocks :: (Matrix a,Matrix a,Matrix a,Matrix a) -> Matrix a+{-# INLINE[1] joinBlocks #-} joinBlocks (tl,tr,bl,br) = (tl <|> tr) <-> (bl <|> br)@@ -505,12 +548,10 @@ -- /This condition is not checked/. (<|>) :: Matrix a -> Matrix a -> Matrix a {-# INLINE (<|>) #-}-(M n m v) <|> (M _ m' v') = M n m'' $ V.generate (n*m'') $- \k -> let (i,j) = decode m'' k in if j <= m- then v V.! encode m (i,j)- else v' V.! encode m' (i,j-m)- where- m'' = m + m'+m <|> m' =+ let c = ncols m+ in matrix (nrows m) (c + ncols m') $ \(i,j) ->+ if j <= c then m ! (i,j) else m' ! (i,j-c) -- | Vertically join two matrices. Visually: --@@ -522,7 +563,10 @@ -- /This condition is not checked/. (<->) :: Matrix a -> Matrix a -> Matrix a {-# INLINE (<->) #-}-(M n m v) <-> (M n' _ v') = M (n+n') m $ v V.++ v'+m <-> m' =+ let r = nrows m+ in matrix (r + nrows m') (ncols m) $ \(i,j) ->+ if i <= r then m ! (i,j) else m' ! (i-r,j) ------------------------------------------------------- -------------------------------------------------------@@ -531,7 +575,8 @@ -- | Perform an operation element-wise. The input matrices are assumed -- to have the same dimensions, but this is not checked. elementwise :: (a -> b -> c) -> (Matrix a -> Matrix b -> Matrix c)-elementwise f (M n m v) (M _ _ v') = M n m $ V.zipWith f v v'+elementwise f m m' = matrix (nrows m) (ncols m) $+ \(i,j) -> f (m ! (i,j)) (m' ! (i,j)) ------------------------------------------------------- -------------------------------------------------------@@ -539,45 +584,107 @@ {- $mult -Three methods are provided for matrix multiplication.+Four methods are provided for matrix multiplication. * 'multStd': Matrix multiplication following directly the definition. This is the best choice when you know for sure that your matrices are small. +* 'multStd2':+ Matrix multiplication following directly the definition.+ However, using a different definition from 'multStd'.+ Some benchmarks show 'multStd2' beats in performance+ to 'multStd' when the size of the matrix is around 100x100.+ * 'multStrassen': Matrix multiplication following the Strassen's algorithm. Complexity grows slower but also some work is added partitioning the matrix. Also, it only works on square matrices of order @2^n@, so if this condition is not met, it is zero-padded until this is accomplished.- Therefore, its use it is not recommended.+ Therefore, its use is not recommended. * 'multStrassenMixed':- This function mixes the 'multStd' and 'multStrassen' methods.+ This function mixes the previous methods. It provides a better performance in general. Method @(@'*'@)@ of the 'Num' class uses this function because it gives the best average performance. However, if you know for sure that your matrices are- small, you should use 'multStd' instead, since- 'multStrassenMixed' is going to switch to that function anyway.+ small (size less than 500x500), you should use 'multStd' or 'multStd2' instead,+ since 'multStrassenMixed' is going to switch to those functions anyway. +We keep researching how to get better performance for matrix multiplication.+If you want to be on the safe side, use ('*').+ -} -- | Standard matrix multiplication by definition. multStd :: Num a => Matrix a -> Matrix a -> Matrix a {-# INLINE multStd #-}-multStd a1@(M n m _) a2@(M n' m' _)+multStd a1@(M n m _ _ _ _) a2@(M n' m' _ _ _ _) -- Checking that sizes match... | m /= n' = error $ "Multiplication of " ++ sizeStr n m ++ " and " ++ sizeStr n' m' ++ " matrices." | otherwise = multStd_ a1 a2 +-- | Standard matrix multiplication by definition.+multStd2 :: Num a => Matrix a -> Matrix a -> Matrix a+{-# INLINE multStd2 #-}+multStd2 a1@(M n m _ _ _ _) a2@(M n' m' _ _ _ _)+ -- Checking that sizes match...+ | m /= n' = error $ "Multiplication of " ++ sizeStr n m ++ " and "+ ++ sizeStr n' m' ++ " matrices."+ | otherwise = multStd__ a1 a2+ -- | Standard matrix multiplication by definition, without checking if sizes match. multStd_ :: Num a => Matrix a -> Matrix a -> Matrix a-{-# INLINE multStd_ #-}-multStd_ a1@(M n m _) a2@(M _ m' _) = matrix n m' $ \(i,j) -> sum [ a1 ! (i,k) * a2 ! (k,j) | k <- [1 .. m] ]+{-# INLINE multStd_ #-}+multStd_ a@(M 1 1 _ _ _ _) b@(M 1 1 _ _ _ _) = M 1 1 0 0 1 $ V.singleton $ (a ! (1,1)) * (b ! (1,1))+multStd_ a@(M 2 2 _ _ _ _) b@(M 2 2 _ _ _ _) =+ M 2 2 0 0 2 $+ let -- A+ a11 = a ! (1,1) ; a12 = a ! (1,2)+ a21 = a ! (2,1) ; a22 = a ! (2,2)+ -- B+ b11 = b ! (1,1) ; b12 = b ! (1,2)+ b21 = b ! (2,1) ; b22 = b ! (2,2)+ in V.fromList + [ a11*b11 + a12*b21 , a11*b12 + a12*b22+ , a21*b11 + a22*b21 , a21*b12 + a22*b22+ ]+multStd_ a@(M 3 3 _ _ _ _) b@(M 3 3 _ _ _ _) =+ M 3 3 0 0 3 $+ let -- A+ a11 = a ! (1,1) ; a12 = a ! (1,2) ; a13 = a ! (1,3)+ a21 = a ! (2,1) ; a22 = a ! (2,2) ; a23 = a ! (2,3)+ a31 = a ! (3,1) ; a32 = a ! (3,2) ; a33 = a ! (3,3)+ -- B+ b11 = b ! (1,1) ; b12 = b ! (1,2) ; b13 = b ! (1,3)+ b21 = b ! (2,1) ; b22 = b ! (2,2) ; b23 = b ! (2,3)+ b31 = b ! (3,1) ; b32 = b ! (3,2) ; b33 = b ! (3,3)+ in V.fromList+ [ a11*b11 + a12*b21 + a13*b31 , a11*b12 + a12*b22 + a13*b32 , a11*b13 + a12*b23 + a13*b33+ , a21*b11 + a22*b21 + a23*b31 , a21*b12 + a22*b22 + a23*b32 , a21*b13 + a22*b23 + a23*b33+ , a31*b11 + a32*b21 + a33*b31 , a31*b12 + a32*b22 + a33*b32 , a31*b13 + a32*b23 + a33*b33+ ]+multStd_ a@(M n m _ _ _ _) b@(M _ m' _ _ _ _) = matrix n m' $ \(i,j) -> sum [ a ! (i,k) * b ! (k,j) | k <- [1 .. m] ] +multStd__ :: Num a => Matrix a -> Matrix a -> Matrix a+{-# INLINE multStd__ #-}+multStd__ a b = matrix r c $ \(i,j) -> dotProduct (V.unsafeIndex avs $ i - 1) (V.unsafeIndex bvs $ j - 1)+ where+ r = nrows a+ avs = V.generate r $ \i -> getRow (i+1) a+ c = ncols b+ bvs = V.generate c $ \i -> getCol (i+1) b++dotProduct :: Num a => V.Vector a -> V.Vector a -> a+{-# INLINE dotProduct #-}+dotProduct v1 v2 = go (V.length v1 - 1) 0+ where+ go (-1) a = a+ go i a = go (i-1) $ (V.unsafeIndex v1 i) * (V.unsafeIndex v2 i) + a+ first :: (a -> Bool) -> [a] -> a first f = go where@@ -587,7 +694,7 @@ -- | Strassen's algorithm over square matrices of order @2^n@. strassen :: Num a => Matrix a -> Matrix a -> Matrix a -- Trivial 1x1 multiplication.-strassen (M 1 1 v) (M 1 1 v') = M 1 1 $ V.zipWith (*) v v'+strassen a@(M 1 1 _ _ _ _) b@(M 1 1 _ _ _ _) = M 1 1 0 0 1 $ V.singleton $ (a ! (1,1)) * (b ! (1,1)) -- General case guesses that the input matrices are square matrices -- whose order is a power of two. strassen a b = joinBlocks (c11,c12,c21,c22)@@ -613,7 +720,7 @@ -- | Strassen's matrix multiplication. multStrassen :: Num a => Matrix a -> Matrix a -> Matrix a-multStrassen a1@(M n m _) a2@(M n' m' _)+multStrassen a1@(M n m _ _ _ _) a2@(M n' m' _ _ _ _) | m /= n' = error $ "Multiplication of " ++ sizeStr n m ++ " and " ++ sizeStr n' m' ++ " matrices." | otherwise =@@ -624,20 +731,22 @@ in submatrix 1 n 1 m' $ strassen b1 b2 strmixFactor :: Int-strmixFactor = 2 ^ (6 :: Int)+strmixFactor = 2 ^ (9 :: Int) -- | Strassen's mixed algorithm. strassenMixed :: Num a => Matrix a -> Matrix a -> Matrix a {-# SPECIALIZE strassenMixed :: Matrix Double -> Matrix Double -> Matrix Double #-} {-# SPECIALIZE strassenMixed :: Matrix Int -> Matrix Int -> Matrix Int #-}-strassenMixed a@(M r _ _) b- | r < strmixFactor = multStd_ a b+{-# SPECIALIZE strassenMixed :: Matrix Rational -> Matrix Rational -> Matrix Rational #-}+strassenMixed a b+ | r < strmixFactor = multStd__ a b | odd r = let r' = r + 1 a' = setSize 0 r' r' a b' = setSize 0 r' r' b in submatrix 1 r 1 r $ strassenMixed a' b' | otherwise = joinBlocks (c11,c12,c21,c22) where+ r = nrows a -- Size of the subproblem is halved. n = quot r 2 -- Split of the original problem into smaller subproblems.@@ -660,7 +769,7 @@ -- | Mixed Strassen's matrix multiplication. multStrassenMixed :: Num a => Matrix a -> Matrix a -> Matrix a {-# INLINE multStrassenMixed #-}-multStrassenMixed a1@(M n m _) a2@(M n' m' _)+multStrassenMixed a1@(M n m _ _ _ _) a2@(M n' m' _ _ _ _) | m /= n' = error $ "Multiplication of " ++ sizeStr n m ++ " and " ++ sizeStr n' m' ++ " matrices." | n < strmixFactor = multStd_ a1 a2@@ -676,7 +785,7 @@ ---- NUMERICAL INSTANCE instance Num a => Num (Matrix a) where- fromInteger = M 1 1 . V.singleton . fromInteger+ fromInteger = M 1 1 0 0 1 . V.singleton . fromInteger negate = fmap negate abs = fmap abs signum = fmap signum@@ -684,22 +793,14 @@ -- Addition of matrices. {-# SPECIALIZE (+) :: Matrix Double -> Matrix Double -> Matrix Double #-} {-# SPECIALIZE (+) :: Matrix Int -> Matrix Int -> Matrix Int #-}- (M n m v) + (M n' m' v')- -- Checking that sizes match...- | n /= n' || m /= m' = error $ "Addition of " ++ sizeStr n m ++ " and "- ++ sizeStr n' m' ++ " matrices."- -- Otherwise, trivial zip.- | otherwise = M n m $ V.zipWith (+) v v'+ {-# SPECIALIZE (+) :: Matrix Rational -> Matrix Rational -> Matrix Rational #-}+ (+) = elementwise (+) -- Substraction of matrices. {-# SPECIALIZE (-) :: Matrix Double -> Matrix Double -> Matrix Double #-} {-# SPECIALIZE (-) :: Matrix Int -> Matrix Int -> Matrix Int #-}- (M n m v) - (M n' m' v')- -- Checking that sizes match...- | n /= n' || m /= m' = error $ "Substraction of " ++ sizeStr n m ++ " and "- ++ sizeStr n' m' ++ " matrices."- -- Otherwise, trivial zip.- | otherwise = M n m $ V.zipWith (-) v v'+ {-# SPECIALIZE (-) :: Matrix Rational -> Matrix Rational -> Matrix Rational #-}+ (-) = elementwise (-) -- Multiplication of matrices. {-# INLINE (*) #-}@@ -746,9 +847,9 @@ -> Int -- ^ Row 2. -> Matrix a -- ^ Original matrix. -> Matrix a -- ^ Matrix with rows 1 and 2 switched.-switchRows r1 r2 (M n m vs) = M n m $ V.modify (\mv -> do+switchRows r1 r2 (M n m ro co w vs) = M n m ro co w $ V.modify (\mv -> do forM_ [1..m] $ \j ->- MV.swap mv (encode m (r1, j)) (encode m (r2, j))) vs+ MV.swap mv (encode w (r1+ro,j+co)) (encode w (r2+ro,j+co))) vs -- | Switch two coumns of a matrix. -- Example:@@ -760,9 +861,9 @@ -> Int -- ^ Col 2. -> Matrix a -- ^ Original matrix. -> Matrix a -- ^ Matrix with cols 1 and 2 switched.-switchCols c1 c2 (M n m vs) = M n m $ V.modify (\mv -> do+switchCols c1 c2 (M n m ro co w vs) = M n m ro co w $ V.modify (\mv -> do forM_ [1..n] $ \j ->- MV.swap mv (encode m (j, c1)) (encode m (j, c2))) vs+ MV.swap mv (encode m (j+ro,c1+co)) (encode m (j+ro,c2+co))) vs ------------------------------------------------------- -------------------------------------------------------@@ -819,11 +920,17 @@ i = maximumBy (\x y -> compare (abs $ u ! (x,k)) (abs $ u ! (y,k))) [ k .. n ] -- Switching to place pivot in current row. u' = switchRows k i u- l' = M (nrows l) (ncols l) $- V.modify (\mv -> mapM_ (\j -> do- msetElem (l ! (k,j)) (ncols l) (i,j) mv- msetElem (l ! (i,j)) (ncols l) (k,j) mv- ) [1 .. k-1] ) $ mvect l+ -- l' = switchRows k i l+ l' = let lw = vcols l+ lro = rowOffset l+ lco = colOffset l+ in if i == k+ then l+ else M (nrows l) (ncols l) lro lco lw $+ V.modify (\mv -> forM_ [1 .. k-1] $ + \j -> MV.swap mv (encode lw (i+lro,j+lco))+ (encode lw (k+lro,j+lco))+ ) $ mvect l p' = switchRows k i p -- Permutation determinant d' = if i == k then d else negate d@@ -902,18 +1009,8 @@ [ (i0, j0) | i0 <- [k .. nrows u], j0 <- [k .. ncols u] ] -- Switching to place pivot in current row. u' = switchCols k j $ switchRows k i u- l'0 = M (nrows l) (ncols l) $- V.modify (\mv -> forM_ [1..k-1] $ \ h -> do- unsafeMset (l ! (k,h)) (ncols l) (i,h) mv- unsafeMset (l ! (i,h)) (ncols l) (k,h) mv- )- $ mvect l- l' = M (nrows l) (ncols l) $- V.modify (\mv -> forM_ [1..k-1] $ \h -> do- unsafeMset (l'0 ! (h,k)) (ncols l) (h,i) mv- unsafeMset (l'0 ! (h,i)) (ncols l) (h,k) mv- )- $ mvect l'0+ l'0 = switchRows k i l+ l' = switchCols k i l'0 p' = switchRows k i p q' = switchCols k j q -- Permutation determinant@@ -951,6 +1048,7 @@ l21 = scaleMatrix (1/l11') a21 a22' = a22 - multStd l21 (transpose l21) l22 = cholDecomp a22'+ ------------------------------------------------------- ------------------------------------------------------- ---- PROPERTIES@@ -996,7 +1094,7 @@ -- consider to use 'detLU' in order to obtain better performance. -- Function 'detLaplace' is /extremely/ slow. detLaplace :: Num a => Matrix a -> a-detLaplace (M 1 1 v) = V.head v+detLaplace m@(M 1 1 _ _ _ _) = m ! (1,1) detLaplace m = sum1 [ (-1)^(i-1) * m ! (i,1) * detLaplace (minorMatrix i 1 m) | i <- [1 .. nrows m] ] where sum1 = foldl1' (+)
bench/mult.hs view
@@ -9,6 +9,9 @@ testdef :: Int -> Matrix Int testdef n = multStd (mat n) (mat n) +testdef2 :: Int -> Matrix Int +testdef2 n = multStd2 (mat n) (mat n) + teststr :: Int -> Matrix Int teststr n = multStrassen (mat n) (mat n) @@ -18,8 +21,10 @@ bmat :: Int -> Benchmark bmat n = bgroup ("mult" ++ show n) [ bench "Definition" $ nf testdef n + , bench "Definition 2" $ nf testdef2 n + -- , bench "Strassen" $ nf teststr n , bench "Strassen mixed" $ nf teststrm n ] main :: IO () -main = defaultMain $ fmap bmat [10,25,100,250,500]+main = defaultMain $ fmap bmat [10,25,100,150,250,500]
matrix.cabal view
@@ -1,5 +1,5 @@ Name: matrix -Version: 0.3.1.1 +Version: 0.3.2.0 Author: Daniel Díaz Category: Math Build-type: Simple
test/Main.hs view
@@ -42,7 +42,6 @@ genMatrix :: Int -> Int -> Gen (Matrix R) genMatrix = genMatrix' - -- | Square matrices newtype Sq = Sq { fromSq :: Matrix R } @@ -58,6 +57,11 @@ main = defaultMain $ testGroup "matrix tests" [ QC.testProperty "identity * m = m * identity = m" $ \(Sq m) -> let n = nrows m in identity n * m == m && m * identity n == m+ , QC.testProperty "a * (b * c) = (a * b) * c"+ $ \(I a) (I b) (I c) (I d) -> forAll (genMatrix a b)+ $ \m1 -> forAll (genMatrix b c)+ $ \m2 -> forAll (genMatrix c d)+ $ \m3 -> m1 * (m2 * m3) == (m1 * m2) * m3 , QC.testProperty "getMatrixAsVector m = mconcat [ getRow i m | i <- [1 .. nrows m]]" $ \m -> getMatrixAsVector (m :: Matrix R) == mconcat [ getRow i m | i <- [1 .. nrows m] ] , QC.testProperty "permMatrix n i j * permMatrix n i j = identity n"@@ -70,6 +74,12 @@ $ \j -> setElem (getElem i j m) (i,j) m == (m :: Matrix R) , QC.testProperty "transpose (transpose m) = m" $ \m -> transpose (transpose m) == (m :: Matrix R)+ , QC.testProperty "if m' = setSize e r c m then (nrows m' = r) && (ncols m' = c)"+ $ \e (I r) (I c) m -> let m' :: Matrix R ; m' = setSize e r c m in nrows m' == r && ncols m' == c+ , QC.testProperty "if (nrows m = r) && (nrcols m = c) then setSize _ r c m = m"+ $ \m -> let r = nrows m+ c = ncols m+ in setSize undefined r c m == (m :: Matrix R) , QC.testProperty "getRow i m = getCol i (transpose m)" $ \m -> forAll (choose (1,nrows m)) $ \i -> getRow i (m :: Matrix R) == getCol i (transpose m)@@ -80,6 +90,12 @@ , QC.testProperty "(+) = elementwise (+)" $ \m1 -> forAll (genMatrix (nrows m1) (ncols m1)) $ \m2 -> m1 + m2 == elementwise (+) m1 m2+ , QC.testProperty "switchCols i j = transpose . switchRows i j . transpose"+ $ \m -> forAll (choose (1,ncols m))+ $ \i -> forAll (choose (1,ncols m))+ $ \j -> switchCols i j (m :: Matrix R) == (transpose $ switchRows i j $ transpose m)+ , QC.testProperty "detLaplace (fromList 3 3 $ repeat 1) = 0"+ $ detLaplace (fromList 3 3 $ repeat 1) == 0 , QC.testProperty "if (u,l,p,d) = luDecomp m then (p*m = l*u) && (detLaplace p = d)" $ \(Sq m) -> (detLaplace m /= 0) ==> (let (u,l,p,d) = luDecompUnsafe m in p*m == l*u && detLaplace p == d)@@ -95,8 +111,4 @@ $ \(Sq m) -> let n = nrows m in forAll (choose (1,n)) $ \i -> forAll (choose (1,n)) $ \j -> detLU (switchRows i j m) == detLU (permMatrix n i j) * detLU m- , QC.testProperty "switchCols i j = transpose . switchRows i j . transpose"- $ \m -> forAll (choose (1,ncols m))- $ \i -> forAll (choose (1,ncols m))- $ \j -> switchCols i j (m :: Matrix R) == (transpose $ switchRows i j $ transpose m) ]