--------------------------------------------------------------------------------
-- |
-- Module      : ArrayFire.Index
-- Copyright   : David Johnson (c) 2019-2026
-- License     : BSD 3
-- Maintainer  : David Johnson <code@dmj.io>
-- Stability   : Experimental
-- Portability : GHC
--
-- Functions for indexing into an 'Array'
--
--------------------------------------------------------------------------------
{-# LANGUAGE FlexibleInstances #-}
module ArrayFire.Index where

import ArrayFire.Internal.Index
import ArrayFire.Internal.Types
import ArrayFire.FFI
import ArrayFire.Exception

import Foreign

import System.IO.Unsafe
import Control.Exception

-- | Index into an 'Array' by 'Seq'
index
  :: Array a
  -- ^ 'Array' argument
  -> [Seq]
  -- ^ 'Seq' to use for indexing
  -> Array a
{-# NOINLINE index #-}
index :: forall a. Array a -> [Seq] -> Array a
index (Array ForeignPtr ()
fptr) [Seq]
seqs =
  IO (Array a) -> Array a
forall a. IO a -> a
unsafePerformIO (IO (Array a) -> Array a)
-> ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> 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) -> IO (Array a))
-> ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a))
-> IO (Array a)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ForeignPtr () -> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. ForeignPtr a -> (Ptr a -> IO b) -> IO b
withForeignPtr ForeignPtr ()
fptr ((AFArray -> IO (Array a)) -> Array a)
-> (AFArray -> IO (Array a)) -> Array a
forall a b. (a -> b) -> a -> b
$ \AFArray
ptr -> do
    (Ptr AFArray -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
alloca ((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
aptr ->
      [AFSeq] -> (Ptr AFSeq -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => [a] -> (Ptr a -> IO b) -> IO b
withArray (Seq -> AFSeq
toAFSeq (Seq -> AFSeq) -> [Seq] -> [AFSeq]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [Seq]
seqs) ((Ptr AFSeq -> IO (Array a)) -> IO (Array a))
-> (Ptr AFSeq -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \Ptr AFSeq
sptr -> do
        AFErr -> IO ()
throwAFError (AFErr -> IO ()) -> IO AFErr -> IO ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Ptr AFArray -> AFArray -> CUInt -> Ptr AFSeq -> IO AFErr
af_index Ptr AFArray
aptr AFArray
ptr CUInt
n Ptr AFSeq
sptr
        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
<$> do
          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
aptr
   where
     n :: CUInt
n = Int -> CUInt
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([Seq] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Seq]
seqs)

-- | Lookup an Array by keys along a specified dimension
lookup
  :: Array a
  -- ^ Input Array
  -> Array Int
  -- ^ Indices
  -> Int
  -- ^ Dimension
  -> Array a
lookup :: forall a. Array a -> Array Int -> Int -> Array a
lookup Array a
a Array Int
b Int
n = Array a
-> Array Int
-> (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
a Array Int
b ((Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a)
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a
forall a b. (a -> b) -> a -> b
$ \Ptr AFArray
p AFArray
x AFArray
y -> Ptr AFArray -> AFArray -> AFArray -> CUInt -> IO AFErr
af_lookup Ptr AFArray
p AFArray
x AFArray
y (Int -> CUInt
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n)

-- | Assign values into an 'Array' range defined by 'Seq' indices
--
-- @
-- >>> let a = vector \@Double 5 [1..]
-- >>> assignSeq a [Seq 1 3 1] (vector \@Double 3 [0,0,0])
-- @
assignSeq
  :: Array a
  -- ^ Destination array
  -> [Seq]
  -- ^ Indices defining the range to assign into
  -> Array a
  -- ^ Source array
  -> Array a
  -- ^ Result with values written at the specified indices
{-# NOINLINE assignSeq #-}
assignSeq :: forall a. Array a -> [Seq] -> Array a -> Array a
assignSeq (Array ForeignPtr ()
fptr) [Seq]
seqs (Array ForeignPtr ()
rhsFptr) =
  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 ()
fptr ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
ptr ->
      ForeignPtr () -> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. ForeignPtr a -> (Ptr a -> IO b) -> IO b
withForeignPtr ForeignPtr ()
rhsFptr ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
rhsPtr ->
        [AFSeq] -> (Ptr AFSeq -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => [a] -> (Ptr a -> IO b) -> IO b
withArray (Seq -> AFSeq
toAFSeq (Seq -> AFSeq) -> [Seq] -> [AFSeq]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [Seq]
seqs) ((Ptr AFSeq -> IO (Array a)) -> IO (Array a))
-> (Ptr AFSeq -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \Ptr AFSeq
sptr ->
          (Ptr AFArray -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
alloca ((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
aptr -> do
            AFErr -> IO ()
throwAFError (AFErr -> IO ()) -> IO AFErr -> IO ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Ptr AFArray -> AFArray -> CUInt -> Ptr AFSeq -> AFArray -> IO AFErr
af_assign_seq Ptr AFArray
aptr AFArray
ptr CUInt
n Ptr AFSeq
sptr AFArray
rhsPtr
            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
aptr)
  where
    n :: CUInt
n = Int -> CUInt
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([Seq] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Seq]
seqs)

-- | Index into an 'Array' using generalized 'Index' values (arrays or sequences)
--
-- @
-- >>> let a = matrix \@Double (3,3) [[1..],[1..],[1..]]
-- >>> indexGen a [seqIdx (Seq 0 1 1) False, seqIdx (Seq 0 1 1) False]
-- @
indexGen
  :: Array a
  -- ^ Input array
  -> [Index]
  -- ^ List of 'Index' values (one per dimension)
  -> Array a
  -- ^ Indexed result
{-# NOINLINE indexGen #-}
indexGen :: forall a. Array a -> [Index] -> Array a
indexGen (Array ForeignPtr ()
fptr) [Index]
indices =
  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 ()
fptr ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
ptr -> do
      afIndices <- (Index -> IO AFIndex) -> [Index] -> IO [AFIndex]
forall (t :: * -> *) (f :: * -> *) a b.
(Traversable t, Applicative f) =>
(a -> f b) -> t a -> f (t b)
forall (f :: * -> *) a b.
Applicative f =>
(a -> f b) -> [a] -> f [b]
traverse Index -> IO AFIndex
toAFIndex [Index]
indices
      withArray afIndices $ \Ptr AFIndex
iptr ->
        (Ptr AFArray -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
alloca ((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
aptr -> do
          AFErr -> IO ()
throwAFError (AFErr -> IO ()) -> IO AFErr -> IO ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Ptr AFArray -> AFArray -> DimT -> Ptr AFIndex -> IO AFErr
af_index_gen Ptr AFArray
aptr AFArray
ptr (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n) Ptr AFIndex
iptr
          (Index -> IO ()) -> [Index] -> IO ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
(a -> m b) -> t a -> m ()
mapM_ Index -> IO ()
touchIdxFPtr [Index]
indices
          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
aptr)
  where
    n :: Int
n = [Index] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Index]
indices
    touchIdxFPtr :: Index -> IO ()
touchIdxFPtr (ArrIndex Bool
_ (Array ForeignPtr ()
p)) = ForeignPtr () -> IO ()
forall a. ForeignPtr a -> IO ()
touchForeignPtr ForeignPtr ()
p
    touchIdxFPtr Index
_ = () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()

-- | Assign values into an 'Array' using generalized 'Index' values
--
-- @
-- >>> let a = matrix \@Double (3,3) [[1..],[1..],[1..]]
-- >>> let b = matrix \@Double (2,2) [[0,0],[0,0]]
-- >>> assignGen a [seqIdx (Seq 0 1 1) False, seqIdx (Seq 0 1 1) False] b
-- @
assignGen
  :: Array a
  -- ^ Destination array
  -> [Index]
  -- ^ List of 'Index' values defining the range to assign into
  -> Array a
  -- ^ Source array
  -> Array a
  -- ^ Result with values written at the specified indices
{-# NOINLINE assignGen #-}
assignGen :: forall a. Array a -> [Index] -> Array a -> Array a
assignGen (Array ForeignPtr ()
fptr) [Index]
indices (Array ForeignPtr ()
rhsFptr) =
  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 ()
fptr ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
ptr ->
      ForeignPtr () -> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. ForeignPtr a -> (Ptr a -> IO b) -> IO b
withForeignPtr ForeignPtr ()
rhsFptr ((AFArray -> IO (Array a)) -> IO (Array a))
-> (AFArray -> IO (Array a)) -> IO (Array a)
forall a b. (a -> b) -> a -> b
$ \AFArray
rhsPtr -> do
        afIndices <- (Index -> IO AFIndex) -> [Index] -> IO [AFIndex]
forall (t :: * -> *) (f :: * -> *) a b.
(Traversable t, Applicative f) =>
(a -> f b) -> t a -> f (t b)
forall (f :: * -> *) a b.
Applicative f =>
(a -> f b) -> [a] -> f [b]
traverse Index -> IO AFIndex
toAFIndex [Index]
indices
        withArray afIndices $ \Ptr AFIndex
iptr ->
          (Ptr AFArray -> IO (Array a)) -> IO (Array a)
forall a b. Storable a => (Ptr a -> IO b) -> IO b
alloca ((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
aptr -> do
            AFErr -> IO ()
throwAFError (AFErr -> IO ()) -> IO AFErr -> IO ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Ptr AFArray
-> AFArray -> DimT -> Ptr AFIndex -> AFArray -> IO AFErr
af_assign_gen Ptr AFArray
aptr AFArray
ptr (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n) Ptr AFIndex
iptr AFArray
rhsPtr
            (Index -> IO ()) -> [Index] -> IO ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
(a -> m b) -> t a -> m ()
mapM_ Index -> IO ()
touchIdxFPtr [Index]
indices
            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
aptr)
  where
    n :: Int
n = [Index] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Index]
indices
    touchIdxFPtr :: Index -> IO ()
touchIdxFPtr (ArrIndex Bool
_ (Array ForeignPtr ()
p)) = ForeignPtr () -> IO ()
forall a. ForeignPtr a -> IO ()
touchForeignPtr ForeignPtr ()
p
    touchIdxFPtr Index
_ = () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()

-- | A special 'Seq' value representing the entire axis of an 'Array'.
-- Hard-coded from include\/af\/seq.h because FFI cannot import static const values.
afSpan :: Seq
afSpan :: Seq
afSpan = Double -> Double -> Double -> Seq
Seq Double
1 Double
1 Double
0

-- | Select the full extent of a dimension. Use in tuple indices where you want all elements along an axis.
--
-- @
-- arr ! (range 0 2, full, at 1)
-- @
full :: Index
full :: Index
full = Bool -> Seq -> Index
SeqIndex Bool
False Seq
afSpan

-- | Convert index expressions to a list of 'Index'.
-- Supports a single 'Index' or tuples of up to four 'Index' values
-- (matching ArrayFire's maximum of 4 dimensions).
class ToIndexList a where
  toIndexList :: a -> [Index]

instance ToIndexList Index where
  toIndexList :: Index -> [Index]
toIndexList Index
x = [Index
x]

instance ToIndexList (Index, Index) where
  toIndexList :: (Index, Index) -> [Index]
toIndexList (Index
a, Index
b) = [Index
a, Index
b]

instance ToIndexList (Index, Index, Index) where
  toIndexList :: (Index, Index, Index) -> [Index]
toIndexList (Index
a, Index
b, Index
c) = [Index
a, Index
b, Index
c]

instance ToIndexList (Index, Index, Index, Index) where
  toIndexList :: (Index, Index, Index, Index) -> [Index]
toIndexList (Index
a, Index
b, Index
c, Index
d) = [Index
a, Index
b, Index
c, Index
d]

-- | Lift a 'Seq' to an 'Index' for use in tuple-based indexing.
idx :: Seq -> Index
idx :: Seq -> Index
idx Seq
s = Bool -> Seq -> Index
SeqIndex Bool
False Seq
s

-- | Index an 'Array'. Accepts a single 'Index' or a tuple of up to four.
--
-- @
-- arr ! at 0                      -- 1D: element 0
-- arr ! range 1 3                 -- 1D: rows 1-3
-- arr ! (range 0 2, at 1)         -- 2D
-- arr ! (range 0 2, full, at 1)   -- 3D, full second axis
-- @
(!) :: ToIndexList ix => Array a -> ix -> Array a
Array a
a ! :: forall ix a. ToIndexList ix => Array a -> ix -> Array a
! ix
ix = Array a -> [Index] -> Array a
forall a. Array a -> [Index] -> Array a
indexGen Array a
a (ix -> [Index]
forall a. ToIndexList a => a -> [Index]
toIndexList ix
ix)
infixl 9 !

-- | Assign into a range of an 'Array'. Lens-style: use with '(&)'.
--
-- @
-- arr & range 1 3 .~ src
-- arr & (range 0 1, at 2) .~ src
-- @
(.~) :: ToIndexList ix => ix -> Array a -> Array a -> Array a
(ix
ix .~ :: forall ix a. ToIndexList ix => ix -> Array a -> Array a -> Array a
.~ Array a
rhs) Array a
arr = Array a -> [Index] -> Array a -> Array a
forall a. Array a -> [Index] -> Array a -> Array a
assignGen Array a
arr (ix -> [Index]
forall a. ToIndexList a => a -> [Index]
toIndexList ix
ix) Array a
rhs
infixr 4 .~