--------------------------------------------------------------------------------
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications    #-}
{-# LANGUAGE ViewPatterns        #-}
--------------------------------------------------------------------------------
-- |
-- Module      : ArrayFire.BLAS
-- Copyright   : David Johnson (c) 2019-2026
-- License     : BSD3
-- Maintainer  : David Johnson <code@dmj.io>
-- Stability   : Experimental
-- Portability : GHC
--
-- Basic Linear Algebra Subprograms (BLAS) API
--
-- @
-- main :: IO ()
-- main = print (matmul x y xProp yProp)
--  where
--     x,y :: Array Double
--     x = matrix (2,3) [[1,2],[3,4],[5,6]]
--     y = matrix (3,2) [[1,2,3],[4,5,6]]
--
--     xProp, yProp :: MatProp
--     xProp = None
--     yProp = None
-- @
-- @
--  ArrayFire Array
--  [2 2 1 1]
--     22.0000    49.0000
--     28.0000    64.0000
-- @
--------------------------------------------------------------------------------
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

-- | The following applies for Sparse-Dense matrix multiplication.
--
-- This function can be used with one sparse input. The sparse input must always be the lhs and the dense matrix must be rhs.
--
-- The sparse array can only be of 'CSR' format.
--
-- The returned array is always dense.
--
-- optLhs an only be one of AF_MAT_NONE, AF_MAT_TRANS, AF_MAT_CTRANS.
--
-- optRhs can only be AF_MAT_NONE.
--
-- >>> matmul (matrix @Double (2,2) [[1,2],[3,4]]) (matrix @Double (2,2) [[1,2],[3,4]]) None None
-- ArrayFire Array
-- [2 2 1 1]
--     7.0000    15.0000
--    10.0000    22.0000
matmul
  :: AFType a
  => Array a
  -- ^ 2D matrix of Array a, left-hand side
  -> Array a
  -- ^ 2D matrix of Array a, right-hand side
  -> MatProp
  -- ^ Left hand side matrix options
  -> MatProp
  -- ^ Right hand side matrix options
  -> Array a
  -- ^ Output of 'matmul'
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))

-- | Plain matrix multiplication — shorthand for @'matmul' a b 'None' 'None'@.
--
-- >>> mm (matrix @Double (2,2) [[1,0],[0,1]]) (matrix @Double (2,2) [[3,4],[5,6]])
-- ArrayFire Array
-- [2 2 1 1]
--     3.0000     5.0000
--     4.0000     6.0000
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

-- | Scalar dot product between two vectors. Also referred to as the inner product.
--
-- >>> dot (vector @Double 10 [1..]) (vector @Double 10 [1..]) None None
-- ArrayFire Array
-- [1 1 1 1]
--   385.0000
dot
  :: AFType a
  => Array a
  -- ^ Left-hand side input
  -> Array a
  -- ^ Right-hand side input
  -> MatProp
  -- ^ Options for left-hand side. Currently only AF_MAT_NONE and AF_MAT_CONJ are supported.
  -> MatProp
  -- ^ Options for right-hand side. Currently only AF_MAT_NONE and AF_MAT_CONJ are supported.
  -> Array a
  -- ^ Output of 'dot'
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))

-- | Scalar dot product between two vectors. Also referred to as the inner product. Returns the result as a host scalar.
--
-- >>> dotAll (vector @Double 10 [1..]) (vector @Double 10 [1..]) None None
-- 385.0
dotAll
  :: forall a . AFResult a
  => Array a
  -- ^ Left-hand side array
  -> Array a
  -- ^ Right-hand side array
  -> MatProp
  -- ^ Options for left-hand side. Currently only AF_MAT_NONE and AF_MAT_CONJ are supported.
  -> MatProp
  -- ^ Options for right-hand side. Currently only AF_MAT_NONE and AF_MAT_CONJ are supported.
  -> Scalar a
  -- ^ Result as the array's element type
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)

-- | Transposes a matrix.
--
-- >>> array = matrix @Double (2,3) [[2,3],[3,4],[5,6]]
-- >>> array
-- ArrayFire Array
-- [2 3 1 1]
--     2.0000     3.0000     5.0000
--     3.0000     4.0000     6.0000
--
-- >>> transpose array True
-- ArrayFire Array
-- [3 2 1 1]
--     2.0000     3.0000
--     3.0000     4.0000
--     5.0000     6.0000
--
transpose
  :: AFType a
  => Array a
  -- ^ Input matrix to be transposed
  -> Bool
  -- ^ Should perform conjugate transposition
  -> Array a
  -- ^ The transposed matrix
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)

-- | Real (non-conjugate) transpose — shorthand for @'transpose' a False@.
--
-- >>> tr (matrix @Double (2,3) [[1,2],[3,4],[5,6]])
-- ArrayFire Array
-- [3 2 1 1]
--     1.0000     3.0000     5.0000
--     2.0000     4.0000     6.0000
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

-- | Transposes a matrix.
--
-- * Warning: This function mutates an array in-place, all subsequent references will be changed. Use carefully.
--
-- >>> array = matrix @Double (2,2) [[1,2],[3,4]]
-- >>> array
-- ArrayFire Array
-- [3 2 1 1]
--    1.0000     4.0000
--    2.0000     5.0000
--    3.0000     6.0000
--
-- >>> transposeInPlace array False
-- >>> array
-- ArrayFire Array
-- [2 2 1 1]
--    1.0000     2.0000
--    3.0000     4.0000
--
transposeInPlace
  :: AFType a
  => Array a
  -- ^ Input matrix to be transposed
  -> Bool
  -- ^ Should perform conjugate transposition
  -> 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)

-- | General Matrix Multiply: C = alpha * op(A) * op(B)
--
-- More general than 'matmul': supports per-element scaling and optional
-- transposition via 'MatProp'.
--
-- >>> gemm None None 1.0 (matrix @Double (2,2) [[1,0],[0,1]]) (matrix @Double (2,2) [[3,4],[5,6]])
-- ArrayFire Array
-- [2 2 1 1]
--     3.0000     5.0000
--     4.0000     6.0000
gemm
  :: forall a . AFType a
  => MatProp
  -- ^ Transformation applied to A ('None', 'Trans', or 'CTrans')
  -> MatProp
  -- ^ Transformation applied to B ('None', 'Trans', or 'CTrans')
  -> a
  -- ^ Scalar alpha
  -> Array a
  -- ^ Matrix A
  -> Array a
  -- ^ Matrix B
  -> Array a
  -- ^ Result C = alpha * op(A) * op(B)
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 #-}