0 | {--
  1 | Copyright (C) 2025  Joel Berkeley
  2 |
  3 | This program is free software: you can redistribute it and/or modify
  4 | it under the terms of the GNU Affero General Public License as published
  5 | by the Free Software Foundation, either version 3 of the License, or
  6 | (at your option) any later version.
  7 |
  8 | This program is distributed in the hope that it will be useful,
  9 | but WITHOUT ANY WARRANTY; without even the implied warranty of
 10 | MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
 11 | GNU Affero General Public License for more details.
 12 |
 13 | You should have received a copy of the GNU Affero General Public License
 14 | along with this program.  If not, see <https://www.gnu.org/licenses/>.
 15 | --}
 16 | ||| For internal spidr use only.
 17 | module Compiler.MLIR.IR.BuiltinAttributes
 18 |
 19 | import Compiler.LLVM.ADT.APFloat
 20 | import Compiler.LLVM.ADT.APInt
 21 | import Compiler.MLIR.IR.Attributes
 22 | import Compiler.MLIR.IR.BuiltinTypes
 23 | import Compiler.MLIR.IR.Types
 24 | import Compiler.Array
 25 | import Compiler.FFI
 26 |
 27 | ffi : String -> String
 28 | ffi = libxla "c/mlir/IR/BuiltinAttributes.h"
 29 |
 30 | public export
 31 | data DenseElementsAttr = MkDenseElementsAttr GCAnyPtr
 32 |
 33 | export
 34 | %foreign (ffi "DenseElementsAttr_delete")
 35 | prim__deleteDenseElementsAttr : AnyPtr -> PrimIO ()
 36 |
 37 | namespace DenseElementsAttr
 38 |   namespace ArrayRef
 39 |     namespace Attribute
 40 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_Attribute")
 41 |       prim__denseElementsAttrGet : GCAnyPtr -> AnyPtr -> Bits64 -> PrimIO AnyPtr
 42 |
 43 |       export
 44 |       get : HasIO io => RankedTensorType -> List Attribute -> io DenseElementsAttr
 45 |       get (MkRankedTensorType rtt) _ = do
 46 |         attr <- primIO $ prim__denseElementsAttrGet rtt prim__getNullAnyPtr 0
 47 |         attr <- onCollectAny' attr (primIO . prim__deleteDenseElementsAttr)
 48 |         pure (MkDenseElementsAttr attr)
 49 |
 50 |     getAux :
 51 |       (GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr) ->
 52 |       HasIO io => RankedTensorType -> Array a -> io DenseElementsAttr
 53 |     getAux f (MkRankedTensorType rtt) (MkArray values len) = do
 54 |       attr <- primIO $ f rtt values len
 55 |       attr <- onCollectAny' attr (primIO . prim__deleteDenseElementsAttr)
 56 |       pure (MkDenseElementsAttr attr)
 57 |
 58 |     namespace Int32
 59 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_int32_t")
 60 |       prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
 61 |
 62 |       export
 63 |       get : HasIO io => RankedTensorType -> Array Int32 -> io DenseElementsAttr
 64 |       get = getAux prim__denseElementsAttrGet
 65 |
 66 |     namespace Int64
 67 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_int64_t")
 68 |       prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
 69 |
 70 |       export
 71 |       get : HasIO io => RankedTensorType -> Array Int64 -> io DenseElementsAttr
 72 |       get = getAux prim__denseElementsAttrGet
 73 |
 74 |     namespace UInt32
 75 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_uint32_t")
 76 |       prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
 77 |
 78 |       export
 79 |       get : HasIO io => RankedTensorType -> Array Bits32 -> io DenseElementsAttr
 80 |       get = getAux prim__denseElementsAttrGet
 81 |
 82 |     namespace UInt64
 83 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_uint64_t")
 84 |       prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
 85 |
 86 |       export
 87 |       get : HasIO io => RankedTensorType -> Array Bits64 -> io DenseElementsAttr
 88 |       get = getAux prim__denseElementsAttrGet
 89 |
 90 |     namespace Double
 91 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_double")
 92 |       prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
 93 |
 94 |       export
 95 |       get : HasIO io => RankedTensorType -> Array Double -> io DenseElementsAttr
 96 |       get = getAux prim__denseElementsAttrGet
 97 |
 98 |     namespace Bool
 99 |       %foreign (ffi "DenseElementsAttr_get_ArrayRef_bool")
100 |       prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> Bits64 -> PrimIO AnyPtr
101 |
102 |       export
103 |       get : HasIO io => RankedTensorType -> Array Bool -> io DenseElementsAttr
104 |       get = getAux prim__denseElementsAttrGet
105 |
106 |   namespace APInt
107 |     %foreign (ffi "DenseElementsAttr_get_APInt")
108 |     prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> PrimIO AnyPtr
109 |
110 |     export
111 |     get : HasIO io => RankedTensorType -> APInt -> io DenseElementsAttr
112 |     get (MkRankedTensorType rtt) (MkAPInt value) = do
113 |       attr <- primIO $ prim__denseElementsAttrGet rtt value
114 |       attr <- onCollectAny' attr (primIO . prim__deleteDenseElementsAttr)
115 |       pure (MkDenseElementsAttr attr)
116 |
117 |   namespace APFloat
118 |     %foreign (ffi "DenseElementsAttr_get_APFloat")
119 |     prim__denseElementsAttrGet : GCAnyPtr -> GCAnyPtr -> PrimIO AnyPtr
120 |
121 |     export
122 |     get : HasIO io => RankedTensorType -> APFloat -> io DenseElementsAttr
123 |     get (MkRankedTensorType rtt) (MkAPFloat value) = do
124 |       attr <- primIO $ prim__denseElementsAttrGet rtt value
125 |       attr <- onCollectAny' attr (primIO . prim__deleteDenseElementsAttr)
126 |       pure (MkDenseElementsAttr attr)
127 |