From b25c499e02fe7e955f4adcc96602d74b98bd89b9 Mon Sep 17 00:00:00 2001 From: Nikita Solodovnikov Date: Thu, 20 Aug 2026 12:24:27 +0300 Subject: [PATCH 1/5] Start implementing differentiation --- ml.cabal | 4 +++ package.yaml | 1 + src/MLambda/Differentiable.hs | 63 +++++++++++++++++++++++++++++++++++ src/MLambda/TypeLits.hs | 32 ++++++++++++++++++ 4 files changed, 100 insertions(+) create mode 100644 src/MLambda/Differentiable.hs diff --git a/ml.cabal b/ml.cabal index 75e726e..feb61b0 100644 --- a/ml.cabal +++ b/ml.cabal @@ -25,6 +25,7 @@ source-repository head library exposed-modules: + MLambda.Differentiable MLambda.Index MLambda.Linear MLambda.Matrix @@ -55,6 +56,7 @@ library , singletons-base , template-haskell , vector + , vinyl default-language: GHC2024 test-suite ml-test @@ -94,6 +96,7 @@ test-suite ml-test , tasty-hunit , template-haskell , vector + , vinyl default-language: GHC2024 benchmark ml-bench @@ -131,4 +134,5 @@ benchmark ml-bench , tasty-bench , template-haskell , vector + , vinyl default-language: GHC2024 diff --git a/package.yaml b/package.yaml index da5d39f..d6f5333 100644 --- a/package.yaml +++ b/package.yaml @@ -32,6 +32,7 @@ dependencies: - singletons - singletons-base - primitive +- vinyl # - ghc-typelits-natnormalise ghc-options: diff --git a/src/MLambda/Differentiable.hs b/src/MLambda/Differentiable.hs new file mode 100644 index 0000000..9679653 --- /dev/null +++ b/src/MLambda/Differentiable.hs @@ -0,0 +1,63 @@ +-- | +-- Module : MLambda.Differentiable +-- Description : Implements differentiation machinery. +-- Copyright : (c) neclitoris, 2026 +-- License : BSD-3-Clause +-- Maintainer : nas140301@gmail.com +-- Stability : experimental +-- Portability : portable +-- +-- This module contains definition of `Differentiable` type class. +{-# LANGUAGE AllowAmbiguousTypes #-} +{-# LANGUAGE PatternSynonyms #-} +module MLambda.Differentiable + ( Functional(..) + , Differentiable(..) + , Matmul(..) + ) where + +import MLambda.Index +import MLambda.Matrix +import MLambda.NDArr +import MLambda.TypeLits + +import Data.Kind +import Data.List.Singletons +import Data.Vinyl +import Numeric.Netlib.Class + +import Prelude hiding (Floating) + + +type family ArgsL (i :: [[Natural]]) e :: [Type] where + ArgsL '[] e = '[] + ArgsL (x ': xs) e = NDArr x e : ArgsL xs e + +type family Args (i :: [[Natural]]) e :: Type where + Args i e = Rec At (Fins @(ArgsL i e) (ArgsL i e)) + +type family Fun (i :: [[Natural]]) (o :: [Natural]) e :: Type where + Fun i o e = Args i e -> NDArr o e + +class Functional f i o e where + ($$) :: f -> Fun i o e + +class Functional f i o e => Differentiable f i o e where + d :: f -> Args i e -> Index o -> Args i e + +data Matmul = Matmul + +instance (KnownNat m, KnownNat n, KnownNat k, Floating e) => Functional Matmul '[[m,n], [n,k]] '[m,k] e where + _ $$ (At a :& At b :& RNil) = a `cross` b + +instance (1 <= m, 1 <= n, 1 <= k, KnownNat m, KnownNat n, KnownNat k, Floating e) => Differentiable Matmul '[[m,n], [n,k]] '[m,k] e where + d _ (At a :& At b :& RNil) (i :. j) = At a' :& At b' :& RNil + where + a' = fromIndex \(k :. l) -> if k == i then b `at` (l :. j) else 0 + b' = fromIndex \(k :. l) -> if l == j then a `at` (i :. k) else 0 + +data (:.:) f1 f2 = f1 :.: f2 + +instance (Functional f1 i1 o1 e, Functional f2 (o1 : i2) o2 e, i ~ i1 ++ i2) => Functional (f1 :.: f2) i o2 e where + (f1 :.: f2) $$ r = undefined + diff --git a/src/MLambda/TypeLits.hs b/src/MLambda/TypeLits.hs index 7f0631d..b9902dd 100644 --- a/src/MLambda/TypeLits.hs +++ b/src/MLambda/TypeLits.hs @@ -1,4 +1,5 @@ {-# LANGUAGE RequiredTypeArguments #-} +{-# LANGUAGE StandaloneKindSignatures #-} -- | -- Module : MLambda.TypeLits @@ -22,11 +23,16 @@ module MLambda.TypeLits , Peano , RNat (..) , RPNat (..) + , Fin (..) + , Fins + , At(..) + , type (!) , ReifiedNat , rnat , rpnat ) where +import Data.Kind import Data.Proxy (Proxy (Proxy)) import GHC.TypeError (ErrorMessage (..), TypeError) import GHC.TypeNats hiding (natVal) @@ -67,6 +73,32 @@ data RPNat n where RPZ :: RPNat PZ RPS :: RPNat n -> RPNat (PS n) +-- | Finite list index. +type Fin :: [k] -> Type +data Fin (l :: [k :: Type]) where + FZ :: Fin (x : xs) + FS :: Fin xs -> Fin (x : xs) + +type (!) :: forall l -> Fin l -> Type +type family (!) (l :: [k :: Type]) (i :: Fin l) where + (x ': _) ! FZ = x + (_ ': xs) ! (FS i) = xs ! i + +type At :: forall {k} {l :: [k]} . Fin l -> Type +data At (i :: Fin l) where + At :: l ! i -> At i + +type Map :: forall k l . (k -> l) -> [k] -> [l] +type family Map f l where + Map _ '[] = '[] + Map f (x ': xs) = f x ': Map f xs + +-- | Creates a list of indices into a type-level list. +type Fins :: forall {k} (l :: [k]) . [k] -> [Fin l] +type family Fins (l :: [k :: Type]) where + Fins '[] = '[] + Fins (x ': xs) = FZ ': Map FS (Fins xs) + -- | A stronger variant of 'KnownNat' which enables induction on type-level naturals. class ReifiedNat n where rnat0 :: RNat n From 059111b52024c23f96eaabde29b0148ffeaf03a7 Mon Sep 17 00:00:00 2001 From: Nikita Solodovnikov Date: Tue, 1 Sep 2026 15:03:55 +0300 Subject: [PATCH 2/5] Implement `Functional` for 'composition' --- src/MLambda/Differentiable.hs | 75 ++++++++++++++++++++++++------ src/MLambda/TypeLits.hs | 87 +++++++++++++++++++++++++++++------ 2 files changed, 135 insertions(+), 27 deletions(-) diff --git a/src/MLambda/Differentiable.hs b/src/MLambda/Differentiable.hs index 9679653..5fd077d 100644 --- a/src/MLambda/Differentiable.hs +++ b/src/MLambda/Differentiable.hs @@ -9,7 +9,7 @@ -- -- This module contains definition of `Differentiable` type class. {-# LANGUAGE AllowAmbiguousTypes #-} -{-# LANGUAGE PatternSynonyms #-} +{-# LANGUAGE RequiredTypeArguments #-} module MLambda.Differentiable ( Functional(..) , Differentiable(..) @@ -21,11 +21,17 @@ import MLambda.Matrix import MLambda.NDArr import MLambda.TypeLits +import Data.Bifunctor +import Data.Either.Singletons import Data.Kind import Data.List.Singletons -import Data.Vinyl +import Data.Singletons +import Data.Type.Equality +import Data.Vinyl hiding ((:~:)) import Numeric.Netlib.Class +import Unsafe.Coerce + import Prelude hiding (Floating) @@ -33,17 +39,15 @@ type family ArgsL (i :: [[Natural]]) e :: [Type] where ArgsL '[] e = '[] ArgsL (x ': xs) e = NDArr x e : ArgsL xs e -type family Args (i :: [[Natural]]) e :: Type where - Args i e = Rec At (Fins @(ArgsL i e) (ArgsL i e)) +type Args i e = Rec At (Fins (ArgsL i e)) -type family Fun (i :: [[Natural]]) (o :: [Natural]) e :: Type where - Fun i o e = Args i e -> NDArr o e +type Fun i o e = Args i e -> NDArr o e class Functional f i o e where ($$) :: f -> Fun i o e class Functional f i o e => Differentiable f i o e where - d :: f -> Args i e -> Index o -> Args i e + d :: f -> Args i e -> Args (Map (Apply (++@#@$) o) i) e data Matmul = Matmul @@ -51,13 +55,56 @@ instance (KnownNat m, KnownNat n, KnownNat k, Floating e) => Functional Matmul ' _ $$ (At a :& At b :& RNil) = a `cross` b instance (1 <= m, 1 <= n, 1 <= k, KnownNat m, KnownNat n, KnownNat k, Floating e) => Differentiable Matmul '[[m,n], [n,k]] '[m,k] e where - d _ (At a :& At b :& RNil) (i :. j) = At a' :& At b' :& RNil - where - a' = fromIndex \(k :. l) -> if k == i then b `at` (l :. j) else 0 - b' = fromIndex \(k :. l) -> if l == j then a `at` (i :. k) else 0 + d _ (At a :& At b :& RNil) = At a' :& At b' :& RNil + where -- TODO: optimize both representation and performance of this + a' = fromIndex \(i :. j :. k :. l) -> if k == i then b `at` (l :. j) else 0 + b' = fromIndex \(i :. j :. k :. l) -> if l == j then a `at` (i :. k) else 0 + +data (:.:) f g = f :.: g + +splitRec :: forall {l :: [Type]} (l1 :: [Type]) (l2 :: [Type]) . + (l ~ l1 ++ l2) => Sing l1 -> Sing l -> Rec At (Fins (l1 ++ l2)) -> (Rec At (Fins l1), Rec At (Fins l2)) +splitRec SNil _ xs = (RNil, xs) +splitRec (SCons _ (sxs :: Sing xs)) (SCons _ ys) (At v :& r) = + first (\lhs -> At v :& shiftFS (withSingI sxs singFins) lhs) $ splitRec @xs @l2 sxs ys (stripFS (withSingI ys singFins) r) + +{-# NOINLINE[1] unFS #-} +{-# RULES "unFSnop" unFS = unsafeCoerce #-} +unFS :: At (FS i) -> At i +unFS (At x) = At x + +{-# NOINLINE[1] doFS #-} +{-# RULES "doFSnop" doFS = unsafeCoerce #-} +doFS :: At i -> At (FS i) +doFS (At x) = At x + +{-# NOINLINE[1] stripFS #-} +{-# RULES "stripFSnop" forall x . stripFS x = unsafeCoerce #-} +stripFS + :: Sing is + -> Rec At (Map (TyCon1 FS) is) + -> Rec At is +stripFS SNil RNil = RNil +stripFS (SCons _ sis) (x :& xs) = unFS x :& stripFS sis xs -data (:.:) f1 f2 = f1 :.: f2 +{-# NOINLINE[1] shiftFS #-} +{-# RULES "shiftFSnop" forall x . shiftFS x = unsafeCoerce #-} +shiftFS + :: Sing is + -> Rec At is + -> Rec At (Map (TyCon1 FS) is) +shiftFS SNil RNil = RNil +shiftFS (SCons _ sis) (x :& xs) = doFS x :& shiftFS sis xs -instance (Functional f1 i1 o1 e, Functional f2 (o1 : i2) o2 e, i ~ i1 ++ i2) => Functional (f1 :.: f2) i o2 e where - (f1 :.: f2) $$ r = undefined +instance + ( Functional f1 i1 o1 e + , Functional f2 (o1 : i2) o2 e + , ArgsL i e ~ ArgsL i1 e ++ ArgsL i2 e, SingI (ArgsL i1 e), SingI (ArgsL i2 e)) + => Functional (f2 :.: f1) i o2 e where + ((f2 :: f2) :.: (f1 :: f1)) $$ r = + let sl1 = sing @(ArgsL i1 e) + sl2 = sing @(ArgsL i2 e) + sl = sl1 %++ sl2 + (lhs, rhs) = splitRec @(ArgsL i1 e) @(ArgsL i2 e) sl1 sl r + in ($$) @f2 @(o1 : i2) @o2 @e f2 (At (($$) @f1 @i1 @o1 @e f1 lhs) :& shiftFS (withSingI sl2 singFins) rhs) diff --git a/src/MLambda/TypeLits.hs b/src/MLambda/TypeLits.hs index b9902dd..d6d657f 100644 --- a/src/MLambda/TypeLits.hs +++ b/src/MLambda/TypeLits.hs @@ -24,8 +24,15 @@ module MLambda.TypeLits , RNat (..) , RPNat (..) , Fin (..) + , SFin (..) + , AppendFin + , headLemma + , appendLaw + , SingFins(..) + , PrependFin , Fins , At(..) + , AtEither(..) , type (!) , ReifiedNat , rnat @@ -33,7 +40,9 @@ module MLambda.TypeLits ) where import Data.Kind -import Data.Proxy (Proxy (Proxy)) +import Data.List.Singletons (Map, SList (..), sMap, type (++)) +import Data.Singletons +import Data.Type.Equality import GHC.TypeError (ErrorMessage (..), TypeError) import GHC.TypeNats hiding (natVal) @@ -75,29 +84,81 @@ data RPNat n where -- | Finite list index. type Fin :: [k] -> Type -data Fin (l :: [k :: Type]) where +data Fin (l :: [k]) where FZ :: Fin (x : xs) FS :: Fin xs -> Fin (x : xs) -type (!) :: forall l -> Fin l -> Type -type family (!) (l :: [k :: Type]) (i :: Fin l) where +type SFin :: forall k (l :: [k]) . Fin l -> Type +data SFin (f :: Fin l) where + SFZ :: SFin FZ + SFS :: forall xs (f :: Fin xs) . SFin f -> SFin (FS f) + +type instance Sing = SFin + +instance SingKind (Fin l) where + type Demote (Fin l) = Fin l + + fromSing SFZ = FZ + fromSing (SFS s) = FS $ fromSing s + + toSing FZ = SomeSing SFZ + toSing (FS s) = (\(SomeSing s') -> SomeSing $ SFS s') $ toSing s + +type (!) :: forall k . forall (l :: [k]) -> Fin l -> k +type family (!) (l :: [k]) (i :: Fin l) where (x ': _) ! FZ = x (_ ': xs) ! (FS i) = xs ! i -type At :: forall {k} {l :: [k]} . Fin l -> Type +type At :: forall {l :: [Type]} . Fin l -> Type data At (i :: Fin l) where - At :: l ! i -> At i - -type Map :: forall k l . (k -> l) -> [k] -> [l] -type family Map f l where - Map _ '[] = '[] - Map f (x ': xs) = f x ': Map f xs + At :: forall {l} (i :: Fin l) . l ! i -> At i + +type AtEither :: forall {l :: [Type]} {r :: [Type]} . Either (Fin l) (Fin r) -> Type +data AtEither (i :: Either (Fin l) (Fin r)) where + AtLeft :: (i ~ Left i') => l ! i' -> AtEither i + AtRight :: (i ~ Right i') => l ! i' -> AtEither i + +type AppendFin :: forall r -> Fin l -> Fin (l ++ r) +type family AppendFin r f where + AppendFin r FZ = FZ + AppendFin r (FS s) = FS (AppendFin r s) + +sAppendFin :: forall r -> forall l (f :: Fin l) . Sing f -> Sing (AppendFin r f) +sAppendFin _ SFZ = SFZ +sAppendFin r (SFS s) = SFS $ sAppendFin r s + +headLemma :: forall x -> ((x:xs) ! FZ) :~: ((x:xs') ! FZ) +headLemma x = trans (Refl :: ((x:xs) ! FZ) :~: x) (Refl :: x :~: ((x:xs') ! FZ)) + +appendLaw :: forall r -> forall l (i :: Fin l) . Sing i -> Sing l -> (l ! i) :~: ((l ++ r) ! AppendFin r i) +appendLaw r si@SFZ (SCons _ _ :: Sing (x:xs)) = + case (sAppendFin r si, si) of + (SFZ :: Sing (AppendFin r i), _ :: Sing i) -> + headLemma @xs @(xs ++ r) x +appendLaw r (SFS s) (SCons _ sxs) = + case appendLaw r s sxs of + Refl -> Refl + +type PrependFin :: forall l -> Fin r -> Fin (l ++ r) +type family PrependFin l f where + PrependFin '[] s = s + PrependFin (x:xs) s = FS (PrependFin xs s) -- | Creates a list of indices into a type-level list. -type Fins :: forall {k} (l :: [k]) . [k] -> [Fin l] +type Fins :: forall {k} . forall (l :: [k]) -> [Fin l] type family Fins (l :: [k :: Type]) where Fins '[] = '[] - Fins (x ': xs) = FZ ': Map FS (Fins xs) + Fins (x ': xs) = FZ ': Map (TyCon1 FS) (Fins xs) + +class SingFins l where + singFins :: Sing (Fins l) + +instance SingI l => SingFins l where + singFins = + case sing @l of + SNil -> SNil + SCons (_ :: Sing x) (sxs :: Sing xs) -> withSingI sxs $ + SCons SFZ (sMap @(Fin xs) @(Fin (x:xs)) @(TyCon1 FS) (SLambda SFS) (singFins @xs)) -- | A stronger variant of 'KnownNat' which enables induction on type-level naturals. class ReifiedNat n where From 51a11667c587aa6f9007f6bc75bc4e26203b38b0 Mon Sep 17 00:00:00 2001 From: Nikita Solodovnikov Date: Tue, 1 Sep 2026 15:04:02 +0300 Subject: [PATCH 3/5] Satisfy linter --- bench/Bench.hs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/bench/Bench.hs b/bench/Bench.hs index 965a3b9..bed68fd 100644 --- a/bench/Bench.hs +++ b/bench/Bench.hs @@ -13,7 +13,7 @@ import Data.Primitive.PrimVar import Data.Random.Normal (normalIO) import Data.Vector.Storable qualified as Storable import Foreign.Storable -import GHC.TypeLits (type (*), type (<=)) +import GHC.TypeLits hiding (natVal) import System.Random (mkStdGen, setStdGen) import Test.Tasty (localOption) import Test.Tasty.Bench @@ -32,7 +32,7 @@ mkNd m n gen = fromIndexM @'[m, n] $ const gen mkVec :: forall m n -> (KnownNat m, KnownNat n, Storable a) => IO a -> IO (Storable.Vector a) -mkVec m n gen = Storable.replicateM (natVal n * natVal m) $ gen +mkVec m n = Storable.replicateM (natVal n * natVal m) mkMassiv :: forall m n -> (KnownNat m, KnownNat n, Storable a) => IO a -> IO (Array S Ix2 a) mkMassiv m n gen = freeze Seq =<< makeMArrayS (Sz2 (natVal n) (natVal m)) (const gen) From 1d6cba07532671fc11c011bb19de6c47322ea4e2 Mon Sep 17 00:00:00 2001 From: Nikita Solodovnikov Date: Tue, 1 Sep 2026 15:11:30 +0300 Subject: [PATCH 4/5] Remove redundant definitions --- src/MLambda/Differentiable.hs | 2 +- src/MLambda/TypeLits.hs | 27 --------------------------- 2 files changed, 1 insertion(+), 28 deletions(-) diff --git a/src/MLambda/Differentiable.hs b/src/MLambda/Differentiable.hs index 5fd077d..2310899 100644 --- a/src/MLambda/Differentiable.hs +++ b/src/MLambda/Differentiable.hs @@ -101,7 +101,7 @@ instance , Functional f2 (o1 : i2) o2 e , ArgsL i e ~ ArgsL i1 e ++ ArgsL i2 e, SingI (ArgsL i1 e), SingI (ArgsL i2 e)) => Functional (f2 :.: f1) i o2 e where - ((f2 :: f2) :.: (f1 :: f1)) $$ r = + (f2 :.: f1) $$ r = let sl1 = sing @(ArgsL i1 e) sl2 = sing @(ArgsL i2 e) sl = sl1 %++ sl2 diff --git a/src/MLambda/TypeLits.hs b/src/MLambda/TypeLits.hs index d6d657f..662b457 100644 --- a/src/MLambda/TypeLits.hs +++ b/src/MLambda/TypeLits.hs @@ -25,14 +25,9 @@ module MLambda.TypeLits , RPNat (..) , Fin (..) , SFin (..) - , AppendFin - , headLemma - , appendLaw , SingFins(..) - , PrependFin , Fins , At(..) - , AtEither(..) , type (!) , ReifiedNat , rnat @@ -42,7 +37,6 @@ module MLambda.TypeLits import Data.Kind import Data.List.Singletons (Map, SList (..), sMap, type (++)) import Data.Singletons -import Data.Type.Equality import GHC.TypeError (ErrorMessage (..), TypeError) import GHC.TypeNats hiding (natVal) @@ -113,32 +107,11 @@ type At :: forall {l :: [Type]} . Fin l -> Type data At (i :: Fin l) where At :: forall {l} (i :: Fin l) . l ! i -> At i -type AtEither :: forall {l :: [Type]} {r :: [Type]} . Either (Fin l) (Fin r) -> Type -data AtEither (i :: Either (Fin l) (Fin r)) where - AtLeft :: (i ~ Left i') => l ! i' -> AtEither i - AtRight :: (i ~ Right i') => l ! i' -> AtEither i - type AppendFin :: forall r -> Fin l -> Fin (l ++ r) type family AppendFin r f where AppendFin r FZ = FZ AppendFin r (FS s) = FS (AppendFin r s) -sAppendFin :: forall r -> forall l (f :: Fin l) . Sing f -> Sing (AppendFin r f) -sAppendFin _ SFZ = SFZ -sAppendFin r (SFS s) = SFS $ sAppendFin r s - -headLemma :: forall x -> ((x:xs) ! FZ) :~: ((x:xs') ! FZ) -headLemma x = trans (Refl :: ((x:xs) ! FZ) :~: x) (Refl :: x :~: ((x:xs') ! FZ)) - -appendLaw :: forall r -> forall l (i :: Fin l) . Sing i -> Sing l -> (l ! i) :~: ((l ++ r) ! AppendFin r i) -appendLaw r si@SFZ (SCons _ _ :: Sing (x:xs)) = - case (sAppendFin r si, si) of - (SFZ :: Sing (AppendFin r i), _ :: Sing i) -> - headLemma @xs @(xs ++ r) x -appendLaw r (SFS s) (SCons _ sxs) = - case appendLaw r s sxs of - Refl -> Refl - type PrependFin :: forall l -> Fin r -> Fin (l ++ r) type family PrependFin l f where PrependFin '[] s = s From 1914396df3f17015e2e357d0f9a8cd4a13c21c00 Mon Sep 17 00:00:00 2001 From: Nikita Solodovnikov Date: Tue, 1 Sep 2026 16:36:03 +0300 Subject: [PATCH 5/5] Fix CI --- ml.cabal | 5 +++-- package.yaml | 3 ++- stack.yaml | 7 ++----- test/Test/MLambda/Matrix.hs | 2 ++ test/Test/MLambda/NDArr.hs | 2 ++ test/Test/MLambda/Utils.hs | 15 +++++++++------ 6 files changed, 20 insertions(+), 14 deletions(-) diff --git a/ml.cabal b/ml.cabal index feb61b0..8a356f9 100644 --- a/ml.cabal +++ b/ml.cabal @@ -1,6 +1,6 @@ cabal-version: 2.2 --- This file has been generated from package.yaml by hpack version 0.38.1. +-- This file has been generated from package.yaml by hpack version 0.39.6. -- -- see: https://github.com/sol/hpack @@ -83,7 +83,7 @@ test-suite ml-test base >=4.7 && <5 , blas-ffi , deepseq - , falsify + , falsify >=0.4.0 && <0.5 , haskell-src-meta , massiv , ml @@ -93,6 +93,7 @@ test-suite ml-test , singletons , singletons-base , tasty + , tasty-falsify , tasty-hunit , template-haskell , vector diff --git a/package.yaml b/package.yaml index d6f5333..9cadee8 100644 --- a/package.yaml +++ b/package.yaml @@ -78,7 +78,8 @@ tests: - ml - tasty - tasty-hunit - - falsify + - falsify >= 0.4.0 && < 0.5 + - tasty-falsify language: GHC2024 default-extensions: diff --git a/stack.yaml b/stack.yaml index 0dd9cb2..c12deed 100644 --- a/stack.yaml +++ b/stack.yaml @@ -17,7 +17,7 @@ # # snapshot: ./custom-snapshot.yaml # snapshot: https://example.com/snapshots/2024-01-01.yaml -snapshot: nightly-2025-08-02 +snapshot: nightly-2026-09-01 # User packages to be built. # Various formats can be used as shown in the example below. @@ -35,7 +35,7 @@ packages: # forks / in-progress versions pinned to a git hash. For example: # extra-deps: -- falsify-0.2.0@sha256:af5c4142095d05775236c8e18e827d540ed22f57b3b74b2268b32040b514ac88,5451 +- tasty-falsify-0.1.0 # # extra-deps: [] @@ -62,6 +62,3 @@ extra-deps: # # Allow a newer minor version of GHC than the snapshot specifies # compiler-check: newer-minor -allow-newer-deps: -- falsify -allow-newer: true diff --git a/test/Test/MLambda/Matrix.hs b/test/Test/MLambda/Matrix.hs index 7b5d525..9558d18 100644 --- a/test/Test/MLambda/Matrix.hs +++ b/test/Test/MLambda/Matrix.hs @@ -1,3 +1,4 @@ +{-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE TypeAbstractions #-} module Test.MLambda.Matrix (testMatrix) where @@ -11,6 +12,7 @@ import Control.Monad import Data.Bool.Singletons import Data.Singletons import GHC.TypeLits.Singletons +import Test.Falsify import Test.Falsify.Predicate ((.$)) import Test.Falsify.Predicate qualified as Pred import Test.Tasty diff --git a/test/Test/MLambda/NDArr.hs b/test/Test/MLambda/NDArr.hs index f5bae98..88bb2ab 100644 --- a/test/Test/MLambda/NDArr.hs +++ b/test/Test/MLambda/NDArr.hs @@ -1,3 +1,4 @@ +{-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE TypeAbstractions #-} {-# LANGUAGE ViewPatterns #-} module Test.MLambda.NDArr (testNDArr) where @@ -15,6 +16,7 @@ import Data.Proxy import Data.Singletons import GHC.TypeLits.Singletons import Prelude.Singletons ((%+)) +import Test.Falsify import Test.Falsify.Generator qualified as Gen import Test.Falsify.Predicate ((.$)) import Test.Falsify.Predicate qualified as Pred diff --git a/test/Test/MLambda/Utils.hs b/test/Test/MLambda/Utils.hs index b3fec3f..6c1e106 100644 --- a/test/Test/MLambda/Utils.hs +++ b/test/Test/MLambda/Utils.hs @@ -5,24 +5,27 @@ import MLambda.Index import MLambda.NDArr import MLambda.TypeLits +import Data.Falsify.ConcreteFun as D +import Data.Falsify.ProperFraction as D import Data.Proxy import Foreign.Storable +import Test.Falsify import Test.Falsify.Generator (Gen) import Test.Falsify.Generator qualified as Gen import Test.Falsify.Range qualified as Range genSz :: Gen Natural -genSz = fromIntegral <$> Gen.int (Range.between (1, 5)) +genSz = fromIntegral <$> Gen.int (Range.inclusive (1, 5)) genDim :: Word -> Word -> Gen [Natural] -genDim a b = Gen.list (Range.between (a, b)) genSz +genDim a b = Gen.list (Range.inclusive (a, b)) genSz genInt :: Gen Int -genInt = Gen.inRange $ Range.between (-1000, 1000) +genInt = Gen.inRange $ Range.inclusive (-1000, 1000) genDouble :: Gen Double genDouble = Gen.inRange $ Range.fromProperFraction 64 - \(Range.ProperFraction d) -> 1 + 4 * d + \(D.ProperFraction d) -> 1 + 4 * d genIndex :: forall dim . Ix dim => Gen (Index dim) genIndex = case inst @dim of @@ -34,9 +37,9 @@ genIndex = case inst @dim of genNDArr :: forall dim e . (Storable e, Ix dim) => Gen e -> Gen (NDArr dim e) genNDArr g = do - Gen.Fn f <- Gen.fun g + Fn f <- Gen.fun g pure $ fromIndex f instance Enum (Index dim) => Gen.Function (Index dim) where - function gb = Gen.functionMap fromEnum toEnum <$> Gen.function gb + function gb = D.map fromEnum toEnum <$> Gen.function gb