packages feed

eigen-hhlo-0.1.0.0: test/EigenHHLO/Test/Cholesky.hs

{-# LANGUAGE DataKinds #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeApplications #-}

module EigenHHLO.Test.Cholesky (tests) where

import qualified Data.Text as T
import Foreign.Ptr (nullPtr)
import Test.Tasty
import Test.Tasty.HUnit

import HHLO.Core.Types (DType(..))
import HHLO.IR.AST (FuncArg(..), TensorType(..))
import HHLO.IR.Builder (arg, moduleFromBuilder)
import HHLO.IR.Pretty (render)
import HHLO.Runtime.PJRT.Types (PJRTApi(..), PJRTClient(..), PJRTDevice(..))
import HHLO.Session (sessionFrom)

import EigenHHLO.IR.Cholesky (cholBuilder)
import EigenHHLO.Core.Types (BackendType(..), EigenSession(..))

dummySess :: EigenSession
dummySess = EigenSession (sessionFrom (PJRTApi nullPtr) (PJRTClient nullPtr) (PJRTDevice nullPtr)) CPU

tests :: TestTree
tests = testGroup "Cholesky"
    [ testCase "MLIR generation" $ do
        let modu = moduleFromBuilder @'[2,2] @'F64 "main"
                [ FuncArg "a" (TensorType [2,2] F64) ]
                $ do a <- arg @'[2,2] @'F64
                     cholBuilder dummySess a
        let mlir = render modu
        assertBool "contains dpotrf" ("eigenhhlo_dpotrf" `T.isInfixOf` mlir)
        assertBool "contains api_version = 2" ("api_version = 2" `T.isInfixOf` mlir)
    ]