{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE ViewPatterns #-}
module ArrayFire.BLAS where
import Control.Exception (mask_)
import Data.Complex
import Foreign.ForeignPtr (newForeignPtr, withForeignPtr)
import Foreign.Marshal.Alloc (alloca)
import Foreign.Marshal.Utils (fillBytes)
import Foreign.Ptr (Ptr, castPtr)
import Foreign.Storable (peek, poke, sizeOf)
import System.IO.Unsafe (unsafePerformIO)
import ArrayFire.Exception
import ArrayFire.FFI
import ArrayFire.Internal.BLAS
import ArrayFire.Internal.Types
matmul
:: AFType a
=> Array a
-> Array a
-> MatProp
-> MatProp
-> Array a
matmul :: forall a.
AFType a =>
Array a -> Array a -> MatProp -> MatProp -> Array a
matmul Array a
arr1 Array a
arr2 MatProp
prop1 MatProp
prop2 = do
Array a
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array a
forall b a c.
Array b
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array c
op2 Array a
arr1 Array a
arr2 (\Ptr AFArray
p AFArray
a AFArray
b -> Ptr AFArray
-> AFArray -> AFArray -> AFMatProp -> AFMatProp -> IO AFErr
af_matmul Ptr AFArray
p AFArray
a AFArray
b (MatProp -> AFMatProp
toMatProp MatProp
prop1) (MatProp -> AFMatProp
toMatProp MatProp
prop2))
mm :: AFType a => Array a -> Array a -> Array a
mm :: forall a. AFType a => Array a -> Array a -> Array a
mm Array a
a Array a
b = Array a -> Array a -> MatProp -> MatProp -> Array a
forall a.
AFType a =>
Array a -> Array a -> MatProp -> MatProp -> Array a
matmul Array a
a Array a
b MatProp
None MatProp
None
dot
:: AFType a
=> Array a
-> Array a
-> MatProp
-> MatProp
-> Array a
dot :: forall a.
AFType a =>
Array a -> Array a -> MatProp -> MatProp -> Array a
dot Array a
arr1 Array a
arr2 MatProp
prop1 MatProp
prop2 =
Array a
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array a
forall b a c.
Array b
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array c
op2 Array a
arr1 Array a
arr2 (\Ptr AFArray
p AFArray
a AFArray
b -> Ptr AFArray
-> AFArray -> AFArray -> AFMatProp -> AFMatProp -> IO AFErr
af_dot Ptr AFArray
p AFArray
a AFArray
b (MatProp -> AFMatProp
toMatProp MatProp
prop1) (MatProp -> AFMatProp
toMatProp MatProp
prop2))
dotAll
:: forall a . AFResult a
=> Array a
-> Array a
-> MatProp
-> MatProp
-> Scalar a
dotAll :: forall a.
AFResult a =>
Array a -> Array a -> MatProp -> MatProp -> Scalar a
dotAll Array a
arr1 Array a
arr2 MatProp
prop1 MatProp
prop2 =
forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a ((Double, Double) -> Scalar a) -> (Double, Double) -> Scalar a
forall a b. (a -> b) -> a -> b
$
Array a
-> Array a
-> (Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr
-> Array arr
-> (Ptr a -> Ptr b -> AFArray -> AFArray -> IO AFErr)
-> (a, b)
infoFromArray22 Array a
arr1 Array a
arr2 ((Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr)
-> (Double, Double))
-> (Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b. (a -> b) -> a -> b
$ \Ptr Double
a Ptr Double
b AFArray
c AFArray
d ->
Ptr Double
-> Ptr Double
-> AFArray
-> AFArray
-> AFMatProp
-> AFMatProp
-> IO AFErr
af_dot_all Ptr Double
a Ptr Double
b AFArray
c AFArray
d (MatProp -> AFMatProp
toMatProp MatProp
prop1) (MatProp -> AFMatProp
toMatProp MatProp
prop2)
transpose
:: AFType a
=> Array a
-> Bool
-> Array a
transpose :: forall a. AFType a => Array a -> Bool -> Array a
transpose Array a
arr1 (Int -> CBool
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> CBool) -> (Bool -> Int) -> Bool -> CBool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Bool -> Int
forall a. Enum a => a -> Int
fromEnum -> CBool
b) =
Array a
arr1 Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array a
forall a b.
Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array b
`op1` (\Ptr AFArray
x AFArray
y -> Ptr AFArray -> AFArray -> CBool -> IO AFErr
af_transpose Ptr AFArray
x AFArray
y CBool
b)
tr :: AFType a => Array a -> Array a
tr :: forall a. AFType a => Array a -> Array a
tr Array a
a = Array a -> Bool -> Array a
forall a. AFType a => Array a -> Bool -> Array a
transpose Array a
a Bool
False
transposeInPlace
:: AFType a
=> Array a
-> Bool
-> IO ()
transposeInPlace :: forall a. AFType a => Array a -> Bool -> IO ()
transposeInPlace Array a
arr (Int -> CBool
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> CBool) -> (Bool -> Int) -> Bool -> CBool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Bool -> Int
forall a. Enum a => a -> Int
fromEnum -> CBool
b) =
Array a
arr Array a -> (AFArray -> IO AFErr) -> IO ()
forall a. Array a -> (AFArray -> IO AFErr) -> IO ()
`inPlace` (AFArray -> CBool -> IO AFErr
`af_transpose_inplace` CBool
b)
gemm
:: forall a . AFType a
=> MatProp
-> MatProp
-> a
-> Array a
-> Array a
-> Array a
gemm :: forall a.
AFType a =>
MatProp -> MatProp -> a -> Array a -> Array a -> Array a
gemm MatProp
opA MatProp
opB a
alpha (Array ForeignPtr ()
fptrA) (Array ForeignPtr ()
fptrB) =
IO (Array a) -> Array a
forall a. IO a -> a
unsafePerformIO (IO (Array a) -> Array a)
-> (IO (Array a) -> IO (Array a)) -> IO (Array a) -> Array a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. IO (Array a) -> IO (Array a)
forall a. IO a -> IO a
mask_ (IO (Array a) -> Array a) -> IO (Array a) -> Array a
forall a b. (a -> b) -> a -> b
$
ForeignPtr () -> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. ForeignPtr a -> (Ptr a -> IO b) -> IO b
withForeignPtr ForeignPtr ()
fptrA ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
ptrA ->
ForeignPtr () -> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. ForeignPtr a -> (Ptr a -> IO b) -> IO b
withForeignPtr ForeignPtr ()
fptrB ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
ptrB ->
(Ptr AFArray -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
calloca ((Ptr AFArray -> IO (Array a)) -> IO (Array a))
-> (Ptr AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \Ptr AFArray
pOut ->
(Ptr a -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
alloca ((Ptr a -> IO (Array a)) -> IO (Array a))
-> (Ptr a -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \Ptr a
pAlpha ->
(Ptr a -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
alloca ((Ptr a -> IO (Array a)) -> IO (Array a))
-> (Ptr a -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \(Ptr a
pBeta :: Ptr a) -> do
Ptr a -> a -> IO ()
forall a. Storable a => Ptr a -> a -> IO ()
poke Ptr a
pAlpha a
alpha
Ptr a -> Word8 -> Int -> IO ()
forall a. Ptr a -> Word8 -> Int -> IO ()
fillBytes Ptr a
pBeta Word8
0 (a -> Int
forall a. Storable a => a -> Int
sizeOf a
alpha)
AFErr -> IO ()
throwAFError (AFErr -> IO ()) -> IO AFErr -> IO ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Ptr AFArray
-> AFMatProp
-> AFMatProp
-> AFArray
-> AFArray
-> AFArray
-> AFArray
-> IO AFErr
af_gemm Ptr AFArray
pOut (MatProp -> AFMatProp
toMatProp MatProp
opA) (MatProp -> AFMatProp
toMatProp MatProp
opB) (Ptr a -> AFArray
forall a b. Ptr a -> Ptr b
castPtr Ptr a
pAlpha) AFArray
ptrA AFArray
ptrB (Ptr a -> AFArray
forall a b. Ptr a -> Ptr b
castPtr Ptr a
pBeta)
ForeignPtr () -> Array a
forall a. ForeignPtr () -> Array a
Array (ForeignPtr () -> Array a) -> IO (ForeignPtr ()) -> IO (Array a)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> (FinalizerPtr () -> AFArray -> IO (ForeignPtr ())
forall a. FinalizerPtr a -> Ptr a -> IO (ForeignPtr a)
newForeignPtr FinalizerPtr ()
af_release_array_finalizer (AFArray -> IO (ForeignPtr ())) -> IO AFArray -> IO (ForeignPtr ())
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Ptr AFArray -> IO AFArray
forall a. Storable a => Ptr a -> IO a
peek Ptr AFArray
pOut)
{-# NOINLINE gemm #-}