17 | module Compiler.LiteralRW
19 | import Control.Monad.State
21 | import Compiler.MLIR.IR.MLIRContext
22 | import Compiler.MLIR.IR.Types
23 | import Compiler.MLIR.IR.BuiltinTypes
24 | import Compiler.MLIR.IR.BuiltinAttributes
25 | import Compiler.Array
26 | import Compiler.DType
33 | dtypeIsArrayType : (dtype : DType) -> ArrayType (idrisType dtype)
34 | dtypeIsArrayType PRED = %search
35 | dtypeIsArrayType S32 = %search
36 | dtypeIsArrayType S64 = %search
37 | dtypeIsArrayType U32 = %search
38 | dtypeIsArrayType U64 = %search
39 | dtypeIsArrayType F64 = %search
42 | mkDenseElementsAttr :
45 | {shape : _} -> {dtype : _} ->
46 | Literal shape (idrisType dtype) ->
47 | io DenseElementsAttr
48 | mkDenseElementsAttr ctx lit = do
49 | type <- RankedTensorType.get shape !(mlirType ctx dtype)
50 | arr <- let at = dtypeIsArrayType dtype in fromList $
toList lit
52 | PRED => DenseElementsAttr.ArrayRef.Bool.get type arr
53 | S32 => DenseElementsAttr.ArrayRef.Int32.get type arr
54 | S64 => DenseElementsAttr.ArrayRef.Int64.get type arr
55 | U32 => DenseElementsAttr.ArrayRef.UInt32.get type arr
56 | U64 => DenseElementsAttr.ArrayRef.UInt64.get type arr
57 | F64 => DenseElementsAttr.ArrayRef.Double.get type arr
60 | read : {shape : _} -> (dtype : DType) -> Array (idrisType dtype) -> Literal shape (idrisType dtype)
62 | let arrIndices = evalState Z $
traverse (\() => State.get <* modify S) (pure ())
63 | at = dtypeIsArrayType dtype
64 | in Array.get arr <$> arrIndices