17 | module Compiler.Stablehlo.Dialect.StablehloAttrs
20 | import Compiler.Array
21 | import Compiler.MLIR.IR.MLIRContext
23 | ffi : String -> String
24 | ffi = libxla "c/stablehlo/dialect/StablehloAttrs.h"
26 | %foreign (ffi "DotDimensionNumbersAttr_delete")
27 | prim__deleteDotDimensionNumbersAttr : AnyPtr -> PrimIO ()
29 | %foreign (ffi "DotDimensionNumbersAttr_get")
30 | prim__dotDimensionNumbersAttrGet :
32 | GCAnyPtr -> Bits64 ->
33 | GCAnyPtr -> Bits64 ->
34 | GCAnyPtr -> Bits64 ->
35 | GCAnyPtr -> Bits64 ->
39 | data DotDimensionNumbersAttr = MkDotDimensionNumbersAttr GCAnyPtr
41 | namespace DotDimensionNumbersAttr
50 | io DotDimensionNumbersAttr
53 | lhsBatchingDimensions
54 | rhsBatchingDimensions
55 | lhsContractingDimensions
56 | rhsContractingDimensions = do
57 | MkArray lb lbLen <- fromList $
cast {to = Int64} <$> lhsBatchingDimensions
58 | MkArray rb rbLen <- fromList $
cast {to = Int64} <$> rhsBatchingDimensions
59 | MkArray lc lcLen <- fromList $
cast {to = Int64} <$> lhsContractingDimensions
60 | MkArray rc rcLen <- fromList $
cast {to = Int64} <$> rhsContractingDimensions
61 | attr <- primIO $
prim__dotDimensionNumbersAttrGet ctx lb lbLen rb rbLen lc lcLen rc rcLen
62 | attr <- onCollectAny' attr (primIO . prim__deleteDotDimensionNumbersAttr)
63 | pure (MkDotDimensionNumbersAttr attr)