packages feed

hhlo-0.6.0.0: test/Test/Runtime/EndToEndAutograd.hs

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

module Test.Runtime.EndToEndAutograd where

import qualified Data.Vector.Storable as V
import Test.Tasty
import Test.Tasty.HUnit

import HHLO.Core.Types
import HHLO.EDSL.Ops
import HHLO.IR.Pretty
import HHLO.Autograd
import HHLO.Runtime.Compile
import HHLO.Runtime.Execute
import HHLO.Runtime.Buffer
import Test.Utils

tests :: TestTree
tests = testGroup "EndToEnd.Autograd"
    [ testCase "grad sum of squares" $ withPJRTCPU $ \api client -> do
        let f x = do sq <- multiply x x; sumAll sq
            modu = gradModule @'[3] @'F32 f
        exec <- compile api client (render modu)
        let inp = V.fromList [1.0, 2.0, 3.0]
        bufIn <- toDeviceF32 api client inp [3]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 3
        -- grad = 2 * x = [2, 4, 6]
        let expected = V.fromList [2.0, 4.0, 6.0]
        assertBool "grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    , testCase "grad sum of doubles" $ withPJRTCPU $ \api client -> do
        let f x = do d <- add x x; sumAll d
            modu = gradModule @'[3] @'F32 f
        exec <- compile api client (render modu)
        let inp = V.fromList [1.0, 2.0, 3.0]
        bufIn <- toDeviceF32 api client inp [3]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 3
        -- grad = [2, 2, 2]
        let expected = V.fromList [2.0, 2.0, 2.0]
        assertBool "grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    , testCase "grad sum of exponentials" $ withPJRTCPU $ \api client -> do
        let f x = do e <- exponential x; sumAll e
            modu = gradModule @'[3] @'F32 f
        exec <- compile api client (render modu)
        let inp = V.fromList [0.0, 0.0, 0.0]
        bufIn <- toDeviceF32 api client inp [3]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 3
        -- grad = exp(x) = [1, 1, 1]
        let expected = V.fromList [1.0, 1.0, 1.0]
        assertBool "grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    , testCase "grad matmul" $ withPJRTCPU $ \api client -> do
        let f x = do
                w <- constant @'[3, 2] @'F32 0.5
                y <- matmul x w
                sumAll y
            gradModu = gradModule @'[2, 3] @'F32 f
        exec <- compile api client (render gradModu)
        let inp = V.fromList [1.0, 2.0, 3.0, 4.0, 5.0, 6.0]
        bufIn <- toDeviceF32 api client inp [2, 3]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 6
        -- grad = sum over cols of W = [0.5+0.5, 0.5+0.5, 0.5+0.5] for each row
        -- = [1.0, 1.0, 1.0, 1.0, 1.0, 1.0]
        let expected = V.fromList [1.0, 1.0, 1.0, 1.0, 1.0, 1.0]
        assertBool "grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    , testCase "grad avgPool" $ withPJRTCPU $ \api client -> do
        let f x = do
                let windowDims = [1, 2, 2, 1]
                    strides    = [1, 2, 2, 1]
                    padding    = replicate 4 [0, 0]
                initVal <- constant @'[] @'F32 0.0
                y <- reduceWindow windowDims strides padding "stablehlo.add" initVal x
                divisor <- constant @'[] @'F32 4.0
                divisorBC <- broadcastWithDims @'[] @'[1, 2, 2, 1] [] divisor
                z <- divide y divisorBC
                sumAll z
            gradModu = gradModule @'[1, 4, 4, 1] @'F32 f
        exec <- compile api client (render gradModu)
        let inp = V.fromList [1.0..16.0]
        bufIn <- toDeviceF32 api client inp [1, 4, 4, 1]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 16
        -- grad = 1/4 for every element (non-overlapping 2x2 avg pool)
        let expected = V.fromList (replicate 16 0.25)
        assertBool "avgPool grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    , testCase "grad conv2d" $ withPJRTCPU $ \api client -> do
        let f x = do
                k <- constant @'[2, 2, 1, 1] @'F32 1.0
                y <- conv2d @1 @3 @3 @1 @1 @2 @2 @2 @2 x k
                sumAll y
            gradModu = gradModule @'[1, 3, 3, 1] @'F32 f
        exec <- compile api client (render gradModu)
        let inp = V.fromList [1.0..9.0]
        bufIn <- toDeviceF32 api client inp [1, 3, 3, 1]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 9
        -- grad for 3x3 input with 2x2 kernel all 1s:
        -- corners: 1, edges: 2, center: 4
        let expected = V.fromList [1, 2, 1, 2, 4, 2, 1, 2, 1]
        assertBool "conv2d grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    , testCase "grad maxPool" $ withPJRTCPU $ \api client -> do
        let f x = do
                let kernel = [2, 2]
                    stride = [2, 2]
                    padding = [[0, 0], [0, 0]]
                y <- maxPool @1 @4 @4 @1 @2 @2 kernel stride padding x
                sumAll y
            gradModu = gradModule @'[1, 4, 4, 1] @'F32 f
        exec <- compile api client (render gradModu)
        let inp = V.fromList [1.0..16.0]
        bufIn <- toDeviceF32 api client inp [1, 4, 4, 1]
        [bufOut] <- execute api exec [bufIn]
        result <- fromDeviceF32 api bufOut 16
        -- maxPool 2x2 stride 2 on 4x4:
        -- Window (0,0): max=6 at pos (1,1) -> index 5
        -- Window (0,1): max=8 at pos (1,3) -> index 7
        -- Window (1,0): max=14 at pos (3,1) -> index 13
        -- Window (1,1): max=16 at pos (3,3) -> index 15
        let expected = V.fromList [0,0,0,0, 0,1,0,1, 0,0,0,0, 0,1,0,1]
        assertBool "maxPool grad close" $
            V.and (V.zipWith (\r e -> abs (r - e) < 0.01) result expected)
    ]