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)
]