hblas 0.2.0.0 → 0.3.0.0
raw patch · 6 files changed
+287/−204 lines, 6 filesPVP ok
version bump matches the API change (PVP)
API changes (from Hackage documentation)
- Numerical.HBLAS.BLAS: coordSwapper :: Transpose -> (a, a) -> (a, a)
- Numerical.HBLAS.BLAS: encodeFFIMatrixHalf :: MatUpLo -> CBLAS_UPLOT
- Numerical.HBLAS.BLAS: encodeFFITranspose :: Transpose -> CBLAS_TRANSPOSET
- Numerical.HBLAS.BLAS: encodeFFITriangleSort :: MatDiag -> CBLAS_DIAGT
- Numerical.HBLAS.BLAS: encodeNiceDIAG :: MatDiag -> BlasDiag
- Numerical.HBLAS.BLAS: encodeNiceOrder :: SOrientation x -> CBLAS_ORDERT
- Numerical.HBLAS.BLAS: encodeNiceTranspose :: Transpose -> BLAS_Transpose
- Numerical.HBLAS.BLAS: encodeNiceUPLO :: MatUpLo -> BLASUplo
- Numerical.HBLAS.BLAS: gemmAbstraction :: (Storable el, PrimMonad m) => String -> GemmFunFFI scale el -> GemmFunFFI scale el -> (el -> (scale -> m ()) -> m ()) -> forall orient. GemmFun el orient (PrimState m) m
- Numerical.HBLAS.BLAS: gemmComplexity :: Integral a => a -> a -> a -> Int64
- Numerical.HBLAS.BLAS: gemvAbstraction :: (Storable el, PrimMonad m) => String -> GemvFunFFI scale el -> GemvFunFFI scale el -> (el -> (scale -> m ()) -> m ()) -> forall orient. GemvFun el orient (PrimState m) m
- Numerical.HBLAS.BLAS: gemvComplexity :: Integral a => a -> a -> Int64
- Numerical.HBLAS.BLAS: isBadGemm :: (Num a, Ord a) => Transpose -> Transpose -> a -> a -> a -> a -> a -> a -> Bool
- Numerical.HBLAS.BLAS: isBadGemv :: Transpose -> Int -> Int -> Int -> Int -> Bool
- Numerical.HBLAS.BLAS: trsvAbstraction :: (Storable el, PrimMonad m) => String -> TrsvFunFFI el -> TrsvFunFFI el -> forall orient. TrsvFun el orient (PrimState m) m
+ Numerical.HBLAS.BLAS.Internal: gemmAbstraction :: (Storable el, PrimMonad m) => String -> GemmFunFFI scale el -> GemmFunFFI scale el -> (el -> (scale -> m ()) -> m ()) -> forall orient. GemmFun el orient (PrimState m) m
+ Numerical.HBLAS.BLAS.Internal: gemvAbstraction :: (Storable el, PrimMonad m) => String -> GemvFunFFI scale el -> GemvFunFFI scale el -> (el -> (scale -> m ()) -> m ()) -> forall orient. GemvFun el orient (PrimState m) m
+ Numerical.HBLAS.BLAS.Internal: trsvAbstraction :: (Storable el, PrimMonad m) => String -> TrsvFunFFI el -> TrsvFunFFI el -> forall orient. TrsvFun el orient (PrimState m) m
- Numerical.HBLAS.BLAS: cgemm :: PrimMonad m => Transpose -> Transpose -> (Complex Float) -> (Complex Float) -> MDenseMatrix (PrimState m) orient (Complex Float) -> MDenseMatrix (PrimState m) orient (Complex Float) -> MDenseMatrix (PrimState m) orient (Complex Float) -> m ()
+ Numerical.HBLAS.BLAS: cgemm :: PrimMonad m => GemmFun (Complex Float) orient (PrimState m) m
- Numerical.HBLAS.BLAS: dgemm :: PrimMonad m => Transpose -> Transpose -> Double -> Double -> MDenseMatrix (PrimState m) orient Double -> MDenseMatrix (PrimState m) orient Double -> MDenseMatrix (PrimState m) orient Double -> m ()
+ Numerical.HBLAS.BLAS: dgemm :: PrimMonad m => GemmFun Double orient (PrimState m) m
- Numerical.HBLAS.BLAS: sgemm :: PrimMonad m => Transpose -> Transpose -> Float -> Float -> MDenseMatrix (PrimState m) orient Float -> MDenseMatrix (PrimState m) orient Float -> MDenseMatrix (PrimState m) orient Float -> m ()
+ Numerical.HBLAS.BLAS: sgemm :: PrimMonad m => GemmFun Float orient (PrimState m) m
- Numerical.HBLAS.BLAS: zgemm :: PrimMonad m => Transpose -> Transpose -> (Complex Double) -> (Complex Double) -> MDenseMatrix (PrimState m) orient (Complex Double) -> MDenseMatrix (PrimState m) orient (Complex Double) -> MDenseMatrix (PrimState m) orient (Complex Double) -> m ()
+ Numerical.HBLAS.BLAS: zgemm :: PrimMonad m => GemmFun (Complex Double) orient (PrimState m) m
Files
- hblas.cabal +3/−2
- src/Numerical/HBLAS/BLAS.hs +68/−193
- src/Numerical/HBLAS/BLAS/Internal.hs +208/−0
- src/Numerical/HBLAS/Constants.lhs +1/−1
- src/Numerical/HBLAS/MatrixTypes.hs +6/−7
- src/Numerical/HBLAS/UtilsFFI.hs +1/−1
hblas.cabal view
@@ -10,7 +10,7 @@ -- PVP summary: +-+------- breaking API changes -- | | +----- non-breaking API additions -- | | | +--- code changes with no API change-version: 0.2.0.0+version: 0.3.0.0 -- A short (one-line) description of the package. synopsis: Human friendly BLAS and Lapack bindings for Haskell.@@ -70,6 +70,7 @@ Numerical.HBLAS.MatrixTypes Numerical.HBLAS.UtilsFFI Numerical.HBLAS.BLAS+ Numerical.HBLAS.BLAS.Internal Numerical.HBLAS.Lapack Numerical.HBLAS.Lapack.FFI Numerical.HBLAS.Constants@@ -111,7 +112,7 @@ -- Base language which the package is written in. default-language: Haskell2010- + ghc-options: -Wall -fno-warn-name-shadowing
src/Numerical/HBLAS/BLAS.hs view
@@ -1,88 +1,89 @@ {-# LANGUAGE BangPatterns , RankNTypes, GADTs, DataKinds #-} +{- | The 'Numerical.HBLAS.BLAS' module provides a fully general+yet type safe BLAS API. -module Numerical.HBLAS.BLAS where---(--- dgemm--- ,sgemm--- ,cgemm--- ,zgemm) +When in doubt about the semantics of an operation,+consult your system's BLAS api documentation, or just read the documentation+for the Intel MKL BLAS distribution +<https://software.intel.com/sites/products/documentation/hpc/mkl/mklman/index.htm the Intel MKL BLAS distribution> -import Numerical.HBLAS.Constants-import Numerical.HBLAS.UtilsFFI -import Numerical.HBLAS.BLAS.FFI -import Numerical.HBLAS.MatrixTypes-import Control.Monad.Primitive-import Data.Complex -import qualified Data.Vector.Storable.Mutable as SM-import Numerical.HBLAS.Constants-import Data.Int +A few basic notes about how to invoke BLAS routines. -gemmComplexity :: Integral a => a -> a -> a -> Int64-gemmComplexity a b c = fromIntegral a * fromIntegral b *fromIntegral c -- this will be wrong by some constant factor, albeit a small one+Many BLAS operations take one or more arguments of type 'Transpose'. +'Tranpose' has the following different constructors, which tell BLAS+routines what transformation to implicitly apply to an input matrix @mat@ with dimension @n x m@. -gemvComplexity :: Integral a => a -> a -> Int64-gemvComplexity a b = fromIntegral a * fromIntegral b +* 'NoTranspose' leaves the matrix @mat@ as is. +* 'Transpose' treats the @mat@ as being implicitly transposed, with dimension+ @m x n@. Entry @mat(i,j)@ being treated as actually being the entry+ @mat(j,i)@. For Real matrices this is also the matrix adjoint operation. + ie @Tranpose(mat)(i,j)=mat(j,i)@ --- this covers the ~6 cases for checking the dimensions for GEMM quite nicely-isBadGemm tra trb ax ay bx by cx cy = isBadGemmHelper (cds tra (ax,ay)) (cds trb (bx,by) ) (cx,cy)- where - cds = coordSwapper - isBadGemmHelper !(ax,ay) !(bx,by) !(cx,cy) = (minimum [ax, ay, bx, by, cx ,cy] <= 0) - || not ( cy == ay && cx == bx && ax == by)+* 'ConjNoTranspose' will implicitly conjugate @mat@, which is a no op for Real ('Float' or 'Double') matrices, but for+'Complex Float' and 'Complex Double' matrices, a given matrix entry @mat(i,j)==x':+'y@+will be treated as actually being @conjugate(mat)(i,j)=y':+'x@. -coordSwapper :: Transpose -> (a,a)-> (a,a)-coordSwapper NoTranspose (a,b) = (a,b)-coordSwapper ConjNoTranspose (a,b) = (a,b) -coordSwapper Transpose (a,b) = (b,a)-coordSwapper ConjTranspose (a,b) = (b,a)+* 'ConjTranpose' will implicitly transpose and conjugate the input matrix.+ConjugateTranpose acts as matrix adjoint for both real and complex matrices. --- / checks if the size of a matrices rows matches input vector size --- and the column count matchesresult vector size-isBadGemv :: Transpose -> Int -> Int -> Int -> Int -> Bool -isBadGemv tr ax ay bdim cdim = isBadGemvHelper (cds tr (ax,ay))- where - cds = coordSwapper- isBadGemvHelper (realX,realY) = - minimum [realY,realX,bdim,cdim] <= 0 || not (realX == bdim && realY == cdim ) -encodeNiceOrder :: SOrientation x -> CBLAS_ORDERT-encodeNiceOrder SRow= encodeOrder BLASRowMajor-encodeNiceOrder SColumn= encodeOrder BLASColMajor+The *gemm operations work as follows (using 'sgemm' as an example): +* @'sgemm trLeft trRight alpha beta left right result'@, where @trLeft@ and @trRight@+are values of type 'Transpose' that respectively act on the matrices @left@ and @right@. -encodeFFITranspose :: Transpose -> CBLAS_TRANSPOSET-encodeFFITranspose x= encodeTranspose $ encodeNiceTranspose x +* the generalized matrix computation thusly formed can be viewed as being +@result = alpha * trLeft(left) * trRight(right) + beta * result@ -encodeNiceTranspose :: Transpose -> BLAS_Transpose-encodeNiceTranspose x = case x of - NoTranspose -> BlasNoTranspose- Transpose -> BlasTranspose- ConjTranspose -> BlasConjTranspose- ConjNoTranspose -> BlasConjNoTranspose -encodeFFIMatrixHalf :: MatUpLo -> CBLAS_UPLOT-encodeFFIMatrixHalf x = encodeUPLO $ encodeNiceUPLO x+the *gemv operations are akin to the *gemm operations, but with @right@ and @result@+being vectors rather than matrices. -encodeNiceUPLO :: MatUpLo -> BLASUplo-encodeNiceUPLO x = case x of- MatUpper -> BUpper- MatLower -> BLower -encodeFFITriangleSort :: MatDiag -> CBLAS_DIAGT-encodeFFITriangleSort x = encodeDiag $ encodeNiceDIAG x+the *trsv operations solve for @x@ in the equation @A x = y@ given @A@ and @y@. +The 'MatUpLo' argument determines if the matrix should be treated as upper or +lower triangular and 'MatDiag' determines if the triangular solver should treat +the diagonal of the matrix as being all 1's or not. A general pattern of invocation+would be @'strsv' matuplo tranposeMatA matdiag matrixA xVector@.+A key detail to note is that the input vector is ALSO the result vector,+ie 'strsv' and friends updates the vector place. -encodeNiceDIAG :: MatDiag -> BlasDiag-encodeNiceDIAG x = case x of- MatUnit -> BlasUnit- MatNonUnit -> BlasNonUnit+-} ---data BLAS_Transpose = BlasNoTranspose | BlasTranspose | BlasConjTranspose | BlasConjNoTranspose ---data Transpose = NoTranspose | Transpose | ConjTranspose | ConjNoTranspose+module Numerical.HBLAS.BLAS(+ GemvFun+ ,GemmFun + ,TrsvFun + + ,dgemm+ ,sgemm+ ,cgemm+ ,zgemm + ,sgemv + ,dgemv+ ,cgemv + ,zgemv + ,strsv+ ,dtrsv+ ,ctrsv+ ,ztrsv + ) where +++import Numerical.HBLAS.UtilsFFI +import Numerical.HBLAS.BLAS.FFI +import Numerical.HBLAS.BLAS.Internal +import Numerical.HBLAS.MatrixTypes+import Control.Monad.Primitive+import Data.Complex +++ type GemmFun el orient s m = Transpose ->Transpose -> el -> el -> MDenseMatrix s orient el -> MDenseMatrix s orient el -> MDenseMatrix s orient el -> m () @@ -95,121 +96,23 @@ -> MDenseMatrix s orient el -> MDenseVector s Direct el -> m () -{--A key design goal of this ffi is to provide *safe* throughput guarantees -for a concurrent application built on top of these apis, while evading-any overheads for providing such safety. Accordingly, on inputs sizes-where the estimated flops count will be more then 1-10 microseconds,-safe ffi calls are used. For inputs whose runtime is under that-unsafe ffi calls are used. --}------- | Matrix mult for general dense matrices---type GemmFunFFI scale el = CBLAS_ORDERT -> CBLAS_TRANSPOSET -> CBLAS_TRANSPOSET->- --CInt -> CInt -> CInt -> {- scal A * B -} scale -> {- Matrix A-} Ptr el -> CInt -> {- B -} Ptr el -> CInt-> - --scale -> {- C -} Ptr el -> CInt -> IO ()---type GemmFun = MutDenseMatrix or el -> MutDenseMatrix or el -> MutDenseMatrix or el -> m ()--{-# NOINLINE gemmAbstraction #-}-gemmAbstraction:: (SM.Storable el, PrimMonad m) => String -> - GemmFunFFI scale el -> GemmFunFFI scale el -> (el -> (scale -> m ())->m ()) -> forall orient . GemmFun el orient (PrimState m) m -gemmAbstraction gemmName gemmSafeFFI gemmUnsafeFFI constHandler = go - where - shouldCallFast :: Int -> Int -> Int -> Bool - shouldCallFast cy cx ax = flopsThreshold >= gemmComplexity cy cx ax-- go tra trb alpha beta - (MutableDenseMatrix ornta ax ay astride abuff) - (MutableDenseMatrix _ bx by bstride bbuff) - (MutableDenseMatrix _ cx cy cstride cbuff) - | isBadGemm tra trb ax ay bx by cx cy = error $! "bad dimension args to GEMM: ax ay bx by cx cy: " ++ show [ax, ay, bx, by, cx ,cy]- | SM.overlaps abuff cbuff || SM.overlaps bbuff cbuff = - error $ "the read and write inputs for: " ++ gemmName ++ " overlap. This is a programmer error. Please fix." - | otherwise = - {- FIXME : Add Sharing check that also errors out for now-}- unsafeWithPrim abuff $ \ap -> - unsafeWithPrim bbuff $ \bp -> - unsafeWithPrim cbuff $ \cp -> - constHandler alpha $ \alphaPtr -> - constHandler beta $ \betaPtr -> - do (ax,ay) <- return $ coordSwapper tra (ax,ay)- --- dont need to swap b, info is in a and c- --- c doesn't get implicitly transposed- blasOrder <- return $ encodeNiceOrder ornta -- all three are the same orientation- rawTra <- return $ encodeFFITranspose tra - rawTrb <- return $ encodeFFITranspose trb- -- example of why i want to switch to singletones- unsafePrimToPrim $! (if shouldCallFast cy cx ax then gemmUnsafeFFI else gemmSafeFFI ) - blasOrder rawTra rawTrb (fromIntegral cy) (fromIntegral cx) (fromIntegral ax) - alphaPtr ap (fromIntegral astride) bp (fromIntegral bstride) betaPtr cp (fromIntegral cstride)---{-pureGemm :: PrimMonad m=>-(Transpose ->Transpose -> el -> el -> MutDenseMatrix (PrimState m) orient el- -> MutDenseMatrix (PrimState m) orient el -> - MutDenseMatrix (PrimState m) orient el -> m ())->- Transpose ->Transpose -> el -> el -> DenseMatrix orient el- -> DenseMatrix orient el -> DenseMatrix orient el -}--sgemm :: PrimMonad m=> - Transpose ->Transpose -> Float -> Float -> MDenseMatrix (PrimState m) orient Float- -> MDenseMatrix (PrimState m) orient Float -> MDenseMatrix (PrimState m) orient Float -> m ()+sgemm :: PrimMonad m=> GemmFun Float orient (PrimState m) m sgemm = gemmAbstraction "sgemm" cblas_sgemm_safe cblas_sgemm_unsafe (\x f -> f x ) -dgemm :: PrimMonad m=> - Transpose ->Transpose -> Double -> Double -> MDenseMatrix (PrimState m) orient Double- -> MDenseMatrix (PrimState m) orient Double -> MDenseMatrix (PrimState m) orient Double -> m ()+dgemm :: PrimMonad m=> GemmFun Double orient (PrimState m) m dgemm = gemmAbstraction "dgemm" cblas_dgemm_safe cblas_dgemm_unsafe (\x f -> f x ) -cgemm :: PrimMonad m=> Transpose ->Transpose -> (Complex Float) -> (Complex Float) -> - MDenseMatrix (PrimState m) orient (Complex Float) -> - MDenseMatrix (PrimState m) orient (Complex Float) -> - MDenseMatrix (PrimState m) orient (Complex Float) -> m ()+cgemm :: PrimMonad m=> GemmFun (Complex Float) orient (PrimState m) m cgemm = gemmAbstraction "cgemm" cblas_cgemm_safe cblas_cgemm_unsafe withRStorable_ -zgemm :: PrimMonad m=> Transpose ->Transpose -> (Complex Double) -> (Complex Double ) -> - MDenseMatrix (PrimState m) orient (Complex Double ) -> - MDenseMatrix (PrimState m) orient (Complex Double) -> - MDenseMatrix (PrimState m) orient (Complex Double) -> m ()+zgemm :: PrimMonad m=> GemmFun (Complex Double) orient (PrimState m) m zgemm = gemmAbstraction "zgemm" cblas_zgemm_safe cblas_zgemm_unsafe withRStorable_ -{-# NOINLINE gemvAbstraction #-}-gemvAbstraction :: (SM.Storable el, PrimMonad m)- => String- -> GemvFunFFI scale el- -> GemvFunFFI scale el- -> (el -> (scale -> m ())-> m ())- -> forall orient . GemvFun el orient (PrimState m) m-gemvAbstraction gemvName gemvSafeFFI gemvUnsafeFFI constHandler = gemv- where- shouldCallFast :: Int -> Int -> Bool- shouldCallFast a b = flopsThreshold >= gemvComplexity a b - gemv tr alpha beta- (MutableDenseMatrix ornta ax ay astride abuff)- (MutableDenseVector _ bdim bstride bbuff)- (MutableDenseVector _ cdim cstride cbuff)- | isBadGemv tr ax ay bdim cdim = error $! "Bad dimension args to GEMV: ax ay xdim ydim: " ++ show [ax, ay, bdim, cdim]- | SM.overlaps abuff cbuff || SM.overlaps bbuff cbuff =- error $! "The read and write inputs for: " ++ gemvName ++ " overlap. This is a programmer error. Please fix." - | otherwise = call- where- (newx,newy) = coordSwapper tr (ax,ay)- call = unsafeWithPrim abuff $ \ap ->- unsafeWithPrim bbuff $ \bp ->- unsafeWithPrim cbuff $ \cp ->- constHandler alpha $ \alphaPtr ->- constHandler beta $ \betaPtr ->- unsafePrimToPrim $! (if shouldCallFast newx newy then gemvUnsafeFFI else gemvSafeFFI)- (encodeNiceOrder ornta) (encodeFFITranspose tr)- (fromIntegral newx) (fromIntegral newy) alphaPtr ap (fromIntegral astride) bp - (fromIntegral bstride) betaPtr cp (fromIntegral cstride)- sgemv :: PrimMonad m => GemvFun Float orient (PrimState m) m sgemv = gemvAbstraction "sgemv" cblas_sgemv_safe cblas_sgemv_unsafe (flip ($)) @@ -221,34 +124,6 @@ zgemv :: PrimMonad m => GemvFun (Complex Double) orient (PrimState m) m zgemv = gemvAbstraction "zgemv" cblas_zgemv_safe cblas_zgemv_unsafe withRStorable_--{-# NOINLINE trsvAbstraction #-}-trsvAbstraction :: (SM.Storable el, PrimMonad m)- => String- -> TrsvFunFFI el -> TrsvFunFFI el- -> forall orient . TrsvFun el orient (PrimState m) m-trsvAbstraction trsvName trsvSafeFFI trsvUnsafeFFI = trsv- where- shouldCallFast :: Int -> Bool- shouldCallFast n = flopsThreshold >= (fromIntegral n)^2-- isBadTrsv :: Int -> Int -> Int -> Bool- isBadTrsv nx ny vdim = nx < 0 || nx /= ny || nx /= vdim-- trsv uplo tra diag- (MutableDenseMatrix ornt x y mstride mbuff)- (MutableDenseVector _ vdim vstride vbuff)- | isBadTrsv x y vdim =- error $! "Bad dimension args to TRSV: x y vdim: " ++ show [x,y,vdim]- | SM.overlaps vbuff mbuff =- error $! "The read and write inputs for: " ++ trsvName ++ " overlap. This is a programmer error. Please fix."- | otherwise = unsafeWithPrim mbuff $ \mp ->- unsafeWithPrim vbuff $ \vp ->- unsafePrimToPrim $! (if shouldCallFast x then trsvUnsafeFFI else trsvSafeFFI)- (encodeNiceOrder ornt) (encodeFFIMatrixHalf uplo) (encodeFFITranspose tra)- (encodeFFITriangleSort diag) (fromIntegral x) mp (fromIntegral mstride) vp- (fromIntegral vstride)- strsv :: PrimMonad m => TrsvFun Float orient (PrimState m) m strsv = trsvAbstraction "strsv" cblas_strsv_safe cblas_strsv_unsafe
+ src/Numerical/HBLAS/BLAS/Internal.hs view
@@ -0,0 +1,208 @@+{-# LANGUAGE BangPatterns , RankNTypes, GADTs, DataKinds #-}++module Numerical.HBLAS.BLAS.Internal(+ gemmAbstraction+ ,gemvAbstraction+ ,trsvAbstraction+ ) where ++import Numerical.HBLAS.Constants+import Numerical.HBLAS.UtilsFFI +import Numerical.HBLAS.BLAS.FFI +import Numerical.HBLAS.MatrixTypes+import Control.Monad.Primitive+import qualified Data.Vector.Storable.Mutable as SM+import Data.Int ++type GemmFun el orient s m = Transpose ->Transpose -> el -> el -> MDenseMatrix s orient el+ -> MDenseMatrix s orient el -> MDenseMatrix s orient el -> m ()++type GemvFun el orient s m = Transpose -> el -> el+ -> MDenseMatrix s orient el -> MDenseVector s Direct el -> MDenseVector s Direct el -> m ()+++type TrsvFun el orient s m =+ MatUpLo -> Transpose -> MatDiag+ -> MDenseMatrix s orient el -> MDenseVector s Direct el -> m () ++gemmComplexity :: Integral a => a -> a -> a -> Int64+gemmComplexity a b c = fromIntegral a * fromIntegral b *fromIntegral c -- this will be wrong by some constant factor, albeit a small one++gemvComplexity :: Integral a => a -> a -> Int64+gemvComplexity a b = fromIntegral a * fromIntegral b +++-- this covers the ~6 cases for checking the dimensions for GEMM quite nicely+isBadGemm :: (Ord a, Num a) =>+ Transpose -> Transpose -> a -> a -> a -> a -> a -> a -> Bool+isBadGemm tra trb ax ay bx by cx cy = isBadGemmHelper (cds tra (ax,ay)) (cds trb (bx,by) ) (cx,cy)+ where + cds = coordSwapper + isBadGemmHelper !(ax,ay) !(bx,by) !(cx,cy) = (minimum [ax, ay, bx, by, cx ,cy] <= 0) + || not ( cy == ay && cx == bx && ax == by)++coordSwapper :: Transpose -> (a,a)-> (a,a)+coordSwapper NoTranspose (a,b) = (a,b)+coordSwapper ConjNoTranspose (a,b) = (a,b) +coordSwapper Transpose (a,b) = (b,a)+coordSwapper ConjTranspose (a,b) = (b,a)++-- / checks if the size of a matrices rows matches input vector size +-- and the column count matchesresult vector size+isBadGemv :: Transpose -> Int -> Int -> Int -> Int -> Bool +isBadGemv tr ax ay bdim cdim = isBadGemvHelper (cds tr (ax,ay))+ where + cds = coordSwapper+ isBadGemvHelper (realX,realY) = + minimum [realY,realX,bdim,cdim] <= 0 || not (realX == bdim && realY == cdim )+++encodeNiceOrder :: SOrientation x -> CBLAS_ORDERT+encodeNiceOrder SRow= encodeOrder BLASRowMajor+encodeNiceOrder SColumn= encodeOrder BLASColMajor+++encodeFFITranspose :: Transpose -> CBLAS_TRANSPOSET+encodeFFITranspose x= encodeTranspose $ encodeNiceTranspose x ++encodeNiceTranspose :: Transpose -> BLAS_Transpose+encodeNiceTranspose x = case x of + NoTranspose -> BlasNoTranspose+ Transpose -> BlasTranspose+ ConjTranspose -> BlasConjTranspose+ ConjNoTranspose -> BlasConjNoTranspose++encodeFFIMatrixHalf :: MatUpLo -> CBLAS_UPLOT+encodeFFIMatrixHalf x = encodeUPLO $ encodeNiceUPLO x++encodeNiceUPLO :: MatUpLo -> BLASUplo+encodeNiceUPLO x = case x of+ MatUpper -> BUpper+ MatLower -> BLower++encodeFFITriangleSort :: MatDiag -> CBLAS_DIAGT+encodeFFITriangleSort x = encodeDiag $ encodeNiceDIAG x++encodeNiceDIAG :: MatDiag -> BlasDiag+encodeNiceDIAG x = case x of+ MatUnit -> BlasUnit+ MatNonUnit -> BlasNonUnit++++++{-+A key design goal of this ffi is to provide *safe* throughput guarantees +for a concurrent application built on top of these apis, while evading+any overheads for providing such safety. Accordingly, on inputs sizes+where the estimated flops count will be more then 1-10 microseconds,+safe ffi calls are used. For inputs whose runtime is under that+unsafe ffi calls are used. +++-}+++---- | Matrix mult for general dense matrices+--type GemmFunFFI scale el = CBLAS_ORDERT -> CBLAS_TRANSPOSET -> CBLAS_TRANSPOSET->+ --CInt -> CInt -> CInt -> {- scal A * B -} scale -> {- Matrix A-} Ptr el -> CInt -> {- B -} Ptr el -> CInt-> + --scale -> {- C -} Ptr el -> CInt -> IO ()+--type GemmFun = MutDenseMatrix or el -> MutDenseMatrix or el -> MutDenseMatrix or el -> m ()++{-# NOINLINE gemmAbstraction #-}+gemmAbstraction:: (SM.Storable el, PrimMonad m) => String -> + GemmFunFFI scale el -> GemmFunFFI scale el -> (el -> (scale -> m ())->m ()) -> forall orient . GemmFun el orient (PrimState m) m +gemmAbstraction gemmName gemmSafeFFI gemmUnsafeFFI constHandler = go + where + shouldCallFast :: Int -> Int -> Int -> Bool + shouldCallFast cy cx ax = flopsThreshold >= gemmComplexity cy cx ax++ go tra trb alpha beta + (MutableDenseMatrix ornta ax ay astride abuff) + (MutableDenseMatrix _ bx by bstride bbuff) + (MutableDenseMatrix _ cx cy cstride cbuff) + | isBadGemm tra trb ax ay bx by cx cy = error $! "bad dimension args to GEMM: ax ay bx by cx cy: " ++ show [ax, ay, bx, by, cx ,cy]+ | SM.overlaps abuff cbuff || SM.overlaps bbuff cbuff = + error $ "the read and write inputs for: " ++ gemmName ++ " overlap. This is a programmer error. Please fix." + | otherwise = + {- FIXME : Add Sharing check that also errors out for now-}+ unsafeWithPrim abuff $ \ap -> + unsafeWithPrim bbuff $ \bp -> + unsafeWithPrim cbuff $ \cp -> + constHandler alpha $ \alphaPtr -> + constHandler beta $ \betaPtr -> + do (axNew,_) <- return $ coordSwapper tra (ax,ay)+ --- dont need to swap b, info is in a and c+ --- c doesn't get implicitly transposed+ blasOrder <- return $ encodeNiceOrder ornta -- all three are the same orientation+ rawTra <- return $ encodeFFITranspose tra + rawTrb <- return $ encodeFFITranspose trb+ -- example of why i want to switch to singletones+ unsafePrimToPrim $! (if shouldCallFast cy cx axNew then gemmUnsafeFFI else gemmSafeFFI ) + blasOrder rawTra rawTrb (fromIntegral cy) (fromIntegral cx) (fromIntegral ax) + alphaPtr ap (fromIntegral astride) bp (fromIntegral bstride) betaPtr cp (fromIntegral cstride)+++++{-# NOINLINE gemvAbstraction #-}+gemvAbstraction :: (SM.Storable el, PrimMonad m)+ => String+ -> GemvFunFFI scale el+ -> GemvFunFFI scale el+ -> (el -> (scale -> m ())-> m ())+ -> forall orient . GemvFun el orient (PrimState m) m+gemvAbstraction gemvName gemvSafeFFI gemvUnsafeFFI constHandler = gemv+ where+ shouldCallFast :: Int -> Int -> Bool+ shouldCallFast a b = flopsThreshold >= gemvComplexity a b + gemv tr alpha beta+ (MutableDenseMatrix ornta ax ay astride abuff)+ (MutableDenseVector _ bdim bstride bbuff)+ (MutableDenseVector _ cdim cstride cbuff)+ | isBadGemv tr ax ay bdim cdim = error $! "Bad dimension args to GEMV: ax ay xdim ydim: " ++ show [ax, ay, bdim, cdim]+ | SM.overlaps abuff cbuff || SM.overlaps bbuff cbuff =+ error $! "The read and write inputs for: " ++ gemvName ++ " overlap. This is a programmer error. Please fix." + | otherwise = call+ where+ (newx,newy) = coordSwapper tr (ax,ay)+ call = unsafeWithPrim abuff $ \ap ->+ unsafeWithPrim bbuff $ \bp ->+ unsafeWithPrim cbuff $ \cp ->+ constHandler alpha $ \alphaPtr ->+ constHandler beta $ \betaPtr ->+ unsafePrimToPrim $! (if shouldCallFast newx newy then gemvUnsafeFFI else gemvSafeFFI)+ (encodeNiceOrder ornta) (encodeFFITranspose tr)+ (fromIntegral newx) (fromIntegral newy) alphaPtr ap (fromIntegral astride) bp + (fromIntegral bstride) betaPtr cp (fromIntegral cstride)++++{-# NOINLINE trsvAbstraction #-}+trsvAbstraction :: (SM.Storable el, PrimMonad m)+ => String+ -> TrsvFunFFI el -> TrsvFunFFI el+ -> forall orient . TrsvFun el orient (PrimState m) m+trsvAbstraction trsvName trsvSafeFFI trsvUnsafeFFI = trsv+ where+ shouldCallFast :: Int -> Bool+ shouldCallFast n = flopsThreshold >= (fromIntegral n :: Int64)^(2 :: Int64)++ isBadTrsv :: Int -> Int -> Int -> Bool+ isBadTrsv nx ny vdim = nx < 0 || nx /= ny || nx /= vdim++ trsv uplo tra diag+ (MutableDenseMatrix ornt x y mstride mbuff)+ (MutableDenseVector _ vdim vstride vbuff)+ | isBadTrsv x y vdim =+ error $! "Bad dimension args to TRSV: x y vdim: " ++ show [x,y,vdim]+ | SM.overlaps vbuff mbuff =+ error $! "The read and write inputs for: " ++ trsvName ++ " overlap. This is a programmer error. Please fix."+ | otherwise = unsafeWithPrim mbuff $ \mp ->+ unsafeWithPrim vbuff $ \vp ->+ unsafePrimToPrim $! (if shouldCallFast x then trsvUnsafeFFI else trsvSafeFFI)+ (encodeNiceOrder ornt) (encodeFFIMatrixHalf uplo) (encodeFFITranspose tra)+ (encodeFFITriangleSort diag) (fromIntegral x) mp (fromIntegral mstride) vp+ (fromIntegral vstride)+
src/Numerical/HBLAS/Constants.lhs view
@@ -1,5 +1,5 @@ \begin{code}-module Numerical.HBLAS.Constants where +module Numerical.HBLAS.Constants(flopsThreshold) where import Data.Int flopsThreshold :: Int64
src/Numerical/HBLAS/MatrixTypes.hs view
@@ -36,11 +36,10 @@ import qualified Data.Vector.Storable as S import qualified Data.Vector.Storable.Mutable as SM import Control.Monad.Primitive ---import Data.Singletons-import Control.Monad.ST.Safe + import Data.Typeable --- import Control.Monad.Primitive + {-@@ -86,13 +85,13 @@ #endif instance Show (SOrientation Row) where- show !a = "SRow"+ show _ = "SRow" instance Show (SOrientation Column) where- show !a = "SColumn" + show _ = "SColumn" instance Eq (SOrientation Row) where- (==) !a !b = True + (==) _ _ = True instance Eq (SOrientation Column) where- (==) !a !b = True + (==) _ _ = True sTranpose :: (x~ TransposeF y, y~TransposeF x ) =>SOrientation x -> SOrientation y
src/Numerical/HBLAS/UtilsFFI.hs view
@@ -9,7 +9,7 @@ import Foreign.ForeignPtr.Safe import Foreign.ForeignPtr.Unsafe -import Foreign.Storable.Complex+import Foreign.Storable.Complex() import Data.Vector.Storable as S import Foreign.Ptr