packages feed

imp-ppl-0.1.0.0: bench/Ellsberg.hs

{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE ExplicitNamespaces #-}
{-# LANGUAGE QualifiedDo #-}
{-# LANGUAGE RebindableSyntax #-}
{-# LANGUAGE UndecidableInstances #-}

-- | An n-way Ellsberg stick-break-chain benchmark.

module Main where

import qualified Imp
import Imp.DSL (Imp, knight, IfThenElse(..))
import Imp.DSL.Combinators (GenNames)
import Imp.DSL.Grade (Merge)

import Prelude
import Data.Proxy (Proxy(..))
import GHC.TypeLits (KnownNat, KnownSymbol, Nat, Symbol, natVal, type (-))

import Bench
  ( BenchCase(..)
  , defaultConfig
  , runAll
  )

type family StickGrade (names :: [Symbol]) :: [Symbol] where
  StickGrade '[]       = '[]
  StickGrade (n ': ns) = Merge '[n] (StickGrade ns)

class StickBreak (names :: [Symbol]) where
  stickChain :: Int -> Imp (StickGrade names) Int

instance StickBreak '[] where
  stickChain = Imp.return

instance (KnownSymbol n, StickBreak ns) => StickBreak (n ': ns) where
  stickChain base = Imp.do
    stop <- knight @n
    if stop
      then Imp.return base
      else stickChain @ns (base + 1)

mkEllsberg
  :: forall (n :: Nat).
     (KnownNat n, StickBreak (GenNames (n - 2) "c"))
  => Imp (StickGrade (GenNames (n - 2) "c")) Int
mkEllsberg = Imp.do
  isZero <- Imp.flip (1 / fromIntegral (natVal (Proxy @n)))
  if isZero
    then Imp.return (0 :: Int)
    else stickChain @(GenNames (n - 2) "c") 1

benchCases :: [BenchCase]
benchCases =
  [ BenchCase 2  (mkEllsberg @2)  (== 1)
  , BenchCase 3  (mkEllsberg @3)  (== 2)
  , BenchCase 4  (mkEllsberg @4)  (== 3)
  , BenchCase 5  (mkEllsberg @5)  (== 4)
  , BenchCase 6  (mkEllsberg @6)  (== 5)
  , BenchCase 7  (mkEllsberg @7)  (== 6)
  , BenchCase 8  (mkEllsberg @8)  (== 7)
  , BenchCase 9  (mkEllsberg @9)  (== 8)
  , BenchCase 10 (mkEllsberg @10) (== 9)
  , BenchCase 11 (mkEllsberg @11) (== 10)
  , BenchCase 12 (mkEllsberg @12) (== 11)
  , BenchCase 13 (mkEllsberg @13) (== 12)
  , BenchCase 14 (mkEllsberg @14) (== 13)
  , BenchCase 15 (mkEllsberg @15) (== 14)
  ]

main :: IO ()
main = runAll defaultConfig benchCases